Nucleus-Resynthesis / runtime /src /resynthesis /intent_scoring_tensor.py
Wl6adams's picture
Add portable Release 188 generation runtime
919fd68 verified
Raw
History Blame Contribute Delete
13.1 kB
"""Tensor-native companion to :mod:`resynthesis.intent_scoring`.
Re-expresses the eight-axis intent rubric as a single batched matmul:
``[batch, 8]`` axis tensor times an ``[8]`` learnable weight vector yields the
composite ``[batch]`` shaping reward. Floors (safety-critical axes) and the
``merits_investigation`` predicate are tensor comparisons -- no Python loops.
Axis order (canonical -- MUST match :data:`AXIS_ORDER`)
------------------------------------------------------
0. chemical_plausibility_completeness
1. distance_from_standard_literature_route
2. starting_material_practicality
3. step_count_route_convergence
4. safety_process_compatibility
5. scalability_impurity_burden
6. evidence_quality_uncertainty
7. merits_investigation
The original :mod:`resynthesis.intent_scoring` module is preserved as the
torch-free reference; this companion is additive and importable independently.
"""
from __future__ import annotations
from collections.abc import Sequence
import torch
from torch import Tensor, nn
from resynthesis.mhc_linear_tensor import MHCExpert
INTENT_SCORING_TENSOR_SCHEMA = "nnf.resynthesis.intent_scoring_tensor.v1"
# Canonical axis order (the columns of the ``[batch, 8]`` input).
AXIS_ORDER: tuple[str, ...] = (
"chemical_plausibility_completeness",
"distance_from_standard_literature_route",
"starting_material_practicality",
"step_count_route_convergence",
"safety_process_compatibility",
"scalability_impurity_burden",
"evidence_quality_uncertainty",
"merits_investigation",
)
NUM_AXES = len(AXIS_ORDER)
# Safety-critical axes (must each clear the reject floor to merit investigation).
CRITICAL_AXIS_INDICES: tuple[int, ...] = (0, 4)
# Default learnable init matches resynthesis.intent_scoring.DEFAULT_INTENT_WEIGHTS.
DEFAULT_INTENT_WEIGHTS: tuple[float, ...] = (
1.5, # plausibility -- chemistry must hold
0.75, # novelty
0.75, # starting material
0.75, # convergence
1.5, # safety -- must hold
0.75, # scalability
1.0, # evidence
0.5, # merit -- preliminary verdict, lightly weighted
)
DEFAULT_INVESTIGATE_THRESHOLD = 0.6
DEFAULT_REJECT_FLOOR = 0.3
RECOMMEND_INVESTIGATE = "investigate"
RECOMMEND_BORDERLINE = "borderline"
RECOMMEND_REJECT = "reject"
# Numeric sentinel for the recommendation tensor (so it flows as a tensor).
REC_INVESTIGATE = 2
REC_BORDERLINE = 1
REC_REJECT = 0
class TensorIntentScorer(nn.Module):
"""Eight-axis intent rubric as a single tensor matmul.
The axis weights ``[8]`` are a learnable :class:`torch.nn.Parameter`
initialized to :data:`DEFAULT_INTENT_WEIGHTS` (matching the Python
reference). The composite is the weight-normalized dot product
``(axes @ weights) / sum(weights)``; ``merits_investigation`` requires the
composite to clear ``investigate_threshold`` AND every safety-critical axis
to clear ``reject_floor``. All ops are differentiable through ``weights``
(and through ``axes`` if the caller makes it require gradients).
MHC-bounded weight residual
---------------------------
When ``mhc_residual=True`` the axis weights are augmented with a bounded
residual produced by an :class:`~resynthesis.mhc_linear_tensor.MHCExpert`
applied to the weight vector itself: the composite becomes
``(axes @ effective_weights) / sum(effective_weights)`` where
``effective_weights = weights + mhc_residual(weights)``. The MHC expert's
``delta`` is bounded to ``[-1, 1]`` via tanh and its ``alpha`` is bounded to
``[0, 1]`` via sigmoid, so the weight adjustment is element-wise bounded --
the intent composite becomes a bounded-residual over the plain weighted
dot product, which keeps the composite's backward gain in a safe band and
enables stable training at a higher learning rate. The default
(``mhc_residual=False``) preserves the original plain-weight behavior for
backward compatibility.
"""
weights: Tensor
mhc_residual_head: MHCExpert | None
def __init__(
self,
*,
weights: Tensor | Sequence[float] | None = None,
investigate_threshold: float = DEFAULT_INVESTIGATE_THRESHOLD,
reject_floor: float = DEFAULT_REJECT_FLOOR,
critical_axes: Sequence[int] = CRITICAL_AXIS_INDICES,
mhc_residual: bool = False,
mhc_sinkhorn_iters: int = 10,
mhc_mix_init: float = 0.9,
device: str | torch.device | None = None,
dtype: torch.dtype = torch.float32,
) -> None:
super().__init__()
if investigate_threshold < 0.0:
raise ValueError("investigate_threshold must be non-negative")
if reject_floor < 0.0:
raise ValueError("reject_floor must be non-negative")
# Normalize device to torch.device | None so downstream constructors
# that require a real device object (not a str) type-check cleanly.
norm_device: torch.device | None = (
torch.device(device) if isinstance(device, str) else device
)
init = (
torch.tensor(DEFAULT_INTENT_WEIGHTS, dtype=dtype, device=device)
if weights is None
else torch.as_tensor(weights, dtype=dtype, device=device)
)
if init.shape != (NUM_AXES,):
raise ValueError(f"weights must have shape ({NUM_AXES},), got {tuple(init.shape)}")
# Full-model construction first instantiates this module on PyTorch's
# meta device, where scalar extraction is unavailable by design. Keep
# the validation tensor-native so meta construction can proceed, while
# concrete CPU/CUDA construction still rejects an invalid supplied
# weight vector before it enters the model.
if init.device.type != "meta":
try:
torch._assert_async(
init.sum() > 0,
"intent weights must sum to a positive value",
)
except (AssertionError, RuntimeError) as error:
raise ValueError(
"intent weights must sum to a positive value"
) from error
self.investigate_threshold = float(investigate_threshold)
self.reject_floor = float(reject_floor)
self.critical_axes: tuple[int, ...] = tuple(critical_axes)
self.dtype = dtype
self.mhc_residual = bool(mhc_residual)
self.weights = nn.Parameter(init.clone())
if self.mhc_residual:
# Bounded-residual expert over the weight vector. Its output is
# element-wise in [-1, 1] (alpha * delta), so the effective weights
# are the plain weights plus a bounded adjustment.
self.mhc_residual_head = MHCExpert(
NUM_AXES,
sinkhorn_iters=mhc_sinkhorn_iters,
mix_init=mhc_mix_init,
dtype=dtype,
device=norm_device,
)
else:
self.mhc_residual_head = None
# ------------------------------------------------------------------
# device helpers
# ------------------------------------------------------------------
@property
def device(self) -> torch.device:
return self.weights.device
def effective_weights(self) -> Tensor:
"""The current per-axis weights, optionally MHC-bounded-residual.
In the default mode this is :attr:`weights` unchanged. In MHC mode
this is ``weights + mhc_residual_head(weights)`` -- the plain weights
plus a bounded ``[-1, 1]`` per-axis residual (alpha * delta) produced
by the MHC expert. Because the residual is bounded, the composite's
backward gain is bounded and the scorer trains stably at a higher LR.
"""
if self.mhc_residual_head is not None:
# The expert expects a trailing feature dim; the weights vector is
# ``[NUM_AXES]`` so we add and remove a leading batch dim.
residual: Tensor = (
self.mhc_residual_head(self.weights.unsqueeze(0)).squeeze(0)
)
return self.weights + residual
return self.weights
def normalized_weights(self) -> Tensor:
"""``effective_weights / sum(effective_weights)`` -- per-axis normalizer.
Uses :meth:`effective_weights` so the MHC-bounded residual (when
enabled) flows through the composite. The sum is clamped away from
zero so a degenerate all-negative residual cannot produce a NaN.
"""
effective = self.effective_weights()
return effective / effective.sum().clamp_min(1e-8)
# ------------------------------------------------------------------
# core scoring (batched, differentiable)
# ------------------------------------------------------------------
def composite(self, axes: Tensor) -> Tensor:
"""Composite shaping reward ``[batch]`` for an ``axes`` ``[batch, 8]``.
``axes`` values should be in ``[0, 1]`` but this method does not clamp --
the caller may want gradients through the raw axis outputs.
"""
normalized = self._check_axes(axes)
return (normalized * self.normalized_weights()).sum(dim=-1)
def floors_pass(self, axes: Tensor) -> Tensor:
"""Boolean ``[batch]``: every safety-critical axis clears ``reject_floor``."""
normalized = self._check_axes(axes)
if not self.critical_axes:
return torch.ones(normalized.shape[0], dtype=torch.bool, device=self.device)
critical = normalized[:, list(self.critical_axes)]
return (critical >= self.reject_floor).all(dim=-1)
def merits_investigation(self, axes: Tensor) -> Tensor:
"""Boolean ``[batch]``: composite clears threshold AND floors pass."""
composite = self.composite(axes)
floors = self.floors_pass(axes)
return (composite >= self.investigate_threshold) & floors
def recommendation_code(self, axes: Tensor) -> Tensor:
"""Numeric recommendation ``[batch]`` (int64): REC_INVESTIGATE/BORDERLINE/REJECT.
Matches :func:`resynthesis.intent_scoring.score_intent`'s thresholds:
investigate if merits + floors; borderline if composite >= threshold*0.7;
else reject.
"""
composite = self.composite(axes)
merits = self.merits_investigation(axes)
borderline_threshold = self.investigate_threshold * 0.7
borderline = (~merits) & (composite >= borderline_threshold)
codes = torch.full_like(merits, REC_REJECT, dtype=torch.long)
codes = torch.where(borderline, torch.full_like(codes, REC_BORDERLINE), codes)
codes = torch.where(merits, torch.full_like(codes, REC_INVESTIGATE), codes)
return codes
def recommendation_str(self, axes: Tensor) -> list[str]:
"""String labels (not differentiable, for logging/inspection)."""
codes = self.recommendation_code(axes).tolist()
lookup = {
REC_INVESTIGATE: RECOMMEND_INVESTIGATE,
REC_BORDERLINE: RECOMMEND_BORDERLINE,
REC_REJECT: RECOMMEND_REJECT,
}
return [lookup[int(c)] for c in codes]
def score(self, axes: Tensor) -> tuple[Tensor, Tensor]:
"""Composite + merits flag ``(composite, merits)`` as tensors."""
return self.composite(axes), self.merits_investigation(axes)
# ------------------------------------------------------------------
# validation
# ------------------------------------------------------------------
def _check_axes(self, axes: Tensor) -> Tensor:
if axes.dim() != 2 or int(axes.shape[-1]) != NUM_AXES:
raise ValueError(
f"axes must have shape [batch, {NUM_AXES}], got {tuple(axes.shape)}"
)
return axes.to(self.device).to(self.dtype)
def axes_tensor(
values: Sequence[Sequence[float]] | Tensor,
*,
dtype: torch.dtype = torch.float32,
device: str | torch.device | None = None,
) -> Tensor:
"""Build a ``[batch, 8]`` axis tensor from a Python sequence."""
if isinstance(values, Tensor):
if values.dim() != 2 or int(values.shape[-1]) != NUM_AXES:
raise ValueError(f"values tensor must be [batch, {NUM_AXES}]")
return values.to(dtype=dtype, device=device)
rows = list(values)
if not rows:
return torch.empty((0, NUM_AXES), dtype=dtype, device=device)
for row in rows:
if len(row) != NUM_AXES:
raise ValueError(
f"each row must have {NUM_AXES} values, got {len(row)}"
)
return torch.tensor(rows, dtype=dtype, device=device)
__all__ = [
"AXIS_ORDER",
"CRITICAL_AXIS_INDICES",
"DEFAULT_INTENT_WEIGHTS",
"DEFAULT_INVESTIGATE_THRESHOLD",
"DEFAULT_REJECT_FLOOR",
"NUM_AXES",
"REC_BORDERLINE",
"REC_INVESTIGATE",
"REC_REJECT",
"RECOMMEND_BORDERLINE",
"RECOMMEND_INVESTIGATE",
"RECOMMEND_REJECT",
"INTENT_SCORING_TENSOR_SCHEMA",
"TensorIntentScorer",
"axes_tensor",
]