| """Tensor-native hard-knowledge surface — trauma as definitive learn targets. |
| |
| **Trauma = hard knowledge (positive + negative).** Every hot-path API here |
| returns ``torch.Tensor`` fields only. JSON / dict receipts live in |
| ``trauma_system.hard_knowledge_surface_receipt`` (explicit boundary adapter). |
| |
| Schema: nnf.resynthesis.hard_knowledge_surface.v1 |
| """ |
|
|
| from __future__ import annotations |
|
|
| from dataclasses import dataclass |
|
|
| import torch |
| from torch import Tensor |
|
|
| from resynthesis.trauma_system import ( |
| DEFAULT_HARD_KNOWLEDGE_FAIL_THRESHOLD, |
| DEFAULT_HARD_WON_EMA_THRESHOLD, |
| DEFAULT_TRAUMA_EPS, |
| DEFAULT_TRAUMA_POSITIVE_SCALE, |
| DEFAULT_TRAUMA_PRESSURE_SCALE, |
| TensorTraumaState, |
| ) |
|
|
| HARD_KNOWLEDGE_SURFACE_SCHEMA = "nnf.resynthesis.hard_knowledge_surface.v1" |
|
|
|
|
| @dataclass(frozen=True) |
| class HardKnowledgeSurfacePacket: |
| """Definitive hard-knowledge surface — all fields are tensors. |
| |
| Negative pole = gaps still missing from weights/pages (must learn). |
| Positive pole = hard-won capability after struggle (must preserve). |
| ``definitive_surface_t`` is the combined router bias (repel gaps, attract wins). |
| """ |
|
|
| schema: str |
| num_arms: int |
| negative_gap_mask_t: Tensor |
| positive_hard_won_mask_t: Tensor |
| negative_gap_signal_t: Tensor |
| positive_hard_won_signal_t: Tensor |
| definitive_surface_t: Tensor |
| empirical_learnability_t: Tensor |
| learning_progress_t: Tensor |
| negative_arm_indices_t: Tensor |
| positive_arm_indices_t: Tensor |
| negative_count_t: Tensor |
| positive_count_t: Tensor |
|
|
|
|
| def _threshold_t(value: float, *, ref: Tensor) -> Tensor: |
| return ref.new_tensor(float(value)) |
|
|
|
|
| def _top_k_arm_indices_t( |
| scores_t: Tensor, |
| mask_t: Tensor, |
| *, |
| top_k: int, |
| ) -> Tensor: |
| """Ranked arm indices with strictly positive masked score (tensor-only).""" |
|
|
| if scores_t.numel() == 0: |
| return scores_t.new_empty(0, dtype=torch.long) |
| masked = scores_t * mask_t.to(dtype=scores_t.dtype) |
| order = masked.argsort(descending=True, stable=True) |
| capped = order[: max(0, int(top_k))] |
| if capped.numel() == 0: |
| return capped.to(dtype=torch.long) |
| keep = masked[capped].gt(0.0) |
| return capped[keep].to(dtype=torch.long) |
|
|
|
|
| def hard_knowledge_surface_packet_t( |
| state: TensorTraumaState, |
| *, |
| top_k: int = 64, |
| negative_threshold: float = DEFAULT_HARD_KNOWLEDGE_FAIL_THRESHOLD, |
| positive_threshold: float = DEFAULT_HARD_WON_EMA_THRESHOLD, |
| negative_scale: float = DEFAULT_TRAUMA_PRESSURE_SCALE, |
| positive_scale: float = DEFAULT_TRAUMA_POSITIVE_SCALE, |
| device: torch.device | None = None, |
| ) -> HardKnowledgeSurfacePacket: |
| """Build the definitive hard-knowledge surface (tensor-native, no host lists).""" |
|
|
| ref = state.fail_ema |
| if device is not None: |
| ref = ref.to(device=device) |
| fails = state.fail_ema.to(device=ref.device, dtype=torch.float32).clamp_min(0.0) |
| peaks = state.peak_fail_ema.to(device=ref.device, dtype=torch.float32).clamp_min(0.0) |
| hard_won = state.hard_won_ema.to(device=ref.device, dtype=torch.float32).clamp_min(0.0) |
| successes = state.success_ema.to(device=ref.device, dtype=torch.float32).clamp_min(0.0) |
|
|
| neg_thr = _threshold_t(negative_threshold, ref=ref) |
| pos_thr = _threshold_t(positive_threshold, ref=ref) |
|
|
| negative_gap_mask_t = ( |
| (fails >= neg_thr) | (peaks >= neg_thr) |
| ).to(dtype=torch.float32) |
| positive_hard_won_mask_t = (hard_won >= pos_thr).to(dtype=torch.float32) |
|
|
| negative_gap_signal_t = (fails + peaks) * negative_gap_mask_t |
| positive_hard_won_signal_t = hard_won * positive_hard_won_mask_t |
|
|
| neg_norm = negative_gap_signal_t.max().clamp_min(DEFAULT_TRAUMA_EPS) |
| pos_norm = positive_hard_won_signal_t.max().clamp_min(DEFAULT_TRAUMA_EPS) |
| neg_unit = negative_gap_signal_t / neg_norm |
| pos_unit = positive_hard_won_signal_t / pos_norm |
| neg_scaled = torch.where( |
| negative_gap_mask_t.gt(0.0), |
| neg_unit * float(negative_scale), |
| neg_unit.new_zeros(()), |
| ) |
| pos_scaled = torch.where( |
| positive_hard_won_mask_t.gt(0.0), |
| pos_unit * float(positive_scale), |
| pos_unit.new_zeros(()), |
| ) |
| |
| |
| |
| definitive_surface_t = pos_scaled - neg_scaled |
|
|
| current_gap_t = fails + peaks |
| empirical_learnability_t = ( |
| successes |
| / (successes + current_gap_t).clamp_min(DEFAULT_TRAUMA_EPS) |
| ).clamp(0.0, 1.0) |
| ownership_mass_t = successes + hard_won |
| learning_progress_t = ( |
| ownership_mass_t |
| / (ownership_mass_t + current_gap_t).clamp_min(DEFAULT_TRAUMA_EPS) |
| ).clamp(0.0, 1.0) |
| mastered_t = positive_hard_won_mask_t.gt(0.0) & negative_gap_mask_t.eq(0.0) |
| learning_progress_t = torch.where( |
| mastered_t, |
| torch.ones_like(learning_progress_t), |
| learning_progress_t, |
| ) |
|
|
| negative_arm_indices_t = _top_k_arm_indices_t( |
| negative_gap_signal_t, |
| negative_gap_mask_t, |
| top_k=top_k, |
| ) |
| positive_arm_indices_t = _top_k_arm_indices_t( |
| positive_hard_won_signal_t, |
| positive_hard_won_mask_t, |
| top_k=top_k, |
| ) |
| negative_count_t = negative_gap_mask_t.sum().to(dtype=torch.long) |
| positive_count_t = positive_hard_won_mask_t.sum().to(dtype=torch.long) |
|
|
| return HardKnowledgeSurfacePacket( |
| schema=HARD_KNOWLEDGE_SURFACE_SCHEMA, |
| num_arms=int(state.num_arms), |
| negative_gap_mask_t=negative_gap_mask_t, |
| positive_hard_won_mask_t=positive_hard_won_mask_t, |
| negative_gap_signal_t=negative_gap_signal_t, |
| positive_hard_won_signal_t=positive_hard_won_signal_t, |
| definitive_surface_t=definitive_surface_t, |
| empirical_learnability_t=empirical_learnability_t, |
| learning_progress_t=learning_progress_t, |
| negative_arm_indices_t=negative_arm_indices_t, |
| positive_arm_indices_t=positive_arm_indices_t, |
| negative_count_t=negative_count_t, |
| positive_count_t=positive_count_t, |
| ) |
|
|
|
|
| def hard_knowledge_router_bias_t( |
| state: TensorTraumaState, |
| *, |
| device: torch.device | None = None, |
| ) -> Tensor: |
| """Per-arm router bias from definitive hard knowledge (tensor-only).""" |
|
|
| return hard_knowledge_surface_packet_t(state, device=device).definitive_surface_t |
|
|
|
|
| def hard_knowledge_mitm_target_levels_t( |
| packet: HardKnowledgeSurfacePacket, |
| *, |
| coverage_floor_t: Tensor | None = None, |
| ) -> Tensor: |
| """Per-arm hardness levels for MITM targeting (negative + coverage floor).""" |
|
|
| levels = packet.negative_gap_signal_t.clone() |
| if coverage_floor_t is not None: |
| floor = coverage_floor_t.to(device=levels.device, dtype=levels.dtype).reshape( |
| -1 |
| ) |
| if floor.numel() == levels.numel(): |
| levels = torch.maximum(levels, floor * packet.negative_gap_mask_t) |
| return levels |
|
|
|
|
| def hard_knowledge_bell_curve_depth_prior_t( |
| *, |
| max_depth: int, |
| center_depth_t: Tensor, |
| scale_depth_t: Tensor, |
| magnitude_t: Tensor, |
| anneal_t: Tensor, |
| device: torch.device | None = None, |
| ) -> Tensor: |
| """Bell-curve depth prior over ``[1, max_depth]`` — tensor-native MITM scaffold.""" |
|
|
| if max_depth < 1: |
| dev = device or center_depth_t.device |
| return torch.empty(0, dtype=torch.float32, device=dev) |
| dev = device or center_depth_t.device |
| axis = torch.arange(1, int(max_depth) + 1, dtype=torch.float32, device=dev) |
| center = center_depth_t.reshape(()).to(device=dev, dtype=torch.float32) |
| scale = scale_depth_t.reshape(()).to(device=dev, dtype=torch.float32).clamp_min( |
| DEFAULT_TRAUMA_EPS |
| ) |
| magnitude = magnitude_t.reshape(()).to(device=dev, dtype=torch.float32).clamp_min(0.0) |
| anneal = anneal_t.reshape(()).to(device=dev, dtype=torch.float32).clamp(0.0, 1.0) |
| z = (axis - center) / scale |
| bell = magnitude * torch.exp(-0.5 * z * z) |
| return torch.nan_to_num(bell * anneal, nan=0.0, posinf=0.0, neginf=0.0) |
|
|
|
|
| __all__ = [ |
| "HARD_KNOWLEDGE_SURFACE_SCHEMA", |
| "HardKnowledgeSurfacePacket", |
| "hard_knowledge_bell_curve_depth_prior_t", |
| "hard_knowledge_mitm_target_levels_t", |
| "hard_knowledge_router_bias_t", |
| "hard_knowledge_surface_packet_t", |
| ] |
|
|