"""Shared building blocks for PNSR and the transcript baselines. The output system (mode -> mechanism -> value) and the cache-record encoder are IDENTICAL across models so that the only manipulated variable is the memory substrate between events (hypothesis §5, §7). """ from __future__ import annotations import math import torch import torch.nn as nn import torch.nn.functional as F from ..world.schema import ENUM_VOCAB, OPS N_MODES, N_ENUM, N_OPS, N_SLOTS = 4, len(ENUM_VOCAB), len(OPS), 100 class SDPA(nn.Module): """Multi-head attention via F.scaled_dot_product_attention (memory- efficient backend; a single broadcastable additive mask, never the per-head materialized weights of nn.MultiheadAttention).""" def __init__(self, d: int, heads: int): super().__init__() self.h, self.hd = heads, d // heads self.q = nn.Linear(d, d, bias=False) self.k = nn.Linear(d, d, bias=False) self.v = nn.Linear(d, d, bias=False) self.o = nn.Linear(d, d, bias=False) def project_kv(self, k, v): B, Tk, _ = k.shape kh = self.k(k).view(B, Tk, self.h, self.hd).transpose(1, 2) vh = self.v(v).view(B, v.shape[1], self.h, self.hd).transpose(1, 2) return kh, vh def forward(self, q, k, v, mask=None, kv_cache=None): B, Tq, _ = q.shape qh = self.q(q).view(B, Tq, self.h, self.hd).transpose(1, 2) kh, vh = kv_cache if kv_cache is not None else self.project_kv(k, v) out = F.scaled_dot_product_attention(qh, kh, vh, attn_mask=mask) return self.o(out.transpose(1, 2).reshape(B, Tq, -1)) def add_mask(pad: torch.Tensor | None, causal: torch.Tensor | None, dtype): """pad: [B,Tk] bool (True = masked) -> [B,1,1,Tk]; causal: [Tq,Tk] bool (True = masked) -> [1,1,Tq,Tk]; summed additive mask or None.""" m = None if pad is not None and pad.any(): m = torch.zeros(pad.shape[0], 1, 1, pad.shape[1], dtype=dtype, device=pad.device).masked_fill(pad[:, None, None, :], float("-inf")) if causal is not None: c = torch.zeros(1, 1, *causal.shape, dtype=dtype, device=causal.device).masked_fill(causal, float("-inf")) m = c if m is None else m + c return m class Block(nn.Module): """Pre-norm transformer block; optional cross-attention.""" def __init__(self, d: int, heads: int, cross: bool = False): super().__init__() self.ln1 = nn.LayerNorm(d) self.attn = SDPA(d, heads) self.cross = cross if cross: self.lnq = nn.LayerNorm(d) self.lnk = nn.LayerNorm(d) self.xattn = SDPA(d, heads) self.ln2 = nn.LayerNorm(d) self.mlp = nn.Sequential(nn.Linear(d, 4 * d), nn.GELU(), nn.Linear(4 * d, d)) def precompute_kv(self, kv): """Project cross-attention K/V once for a KV set that stays constant across parameter-shared deliberation steps.""" kvn = self.lnk(kv) return self.xattn.project_kv(kvn, kvn) def forward(self, x, kv=None, key_padding_mask=None, attn_mask=None, self_padding_mask=None, kv_cache=None): h = self.ln1(x) m = add_mask(self_padding_mask, attn_mask, x.dtype) x = x + self.attn(h, h, h, mask=m) if self.cross and (kv is not None or kv_cache is not None): mk = add_mask(key_padding_mask, None, x.dtype) if kv_cache is not None: x = x + self.xattn(self.lnq(x), None, None, mask=mk, kv_cache=kv_cache) else: kvn = self.lnk(kv) x = x + self.xattn(self.lnq(x), kvn, kvn, mask=mk) return x + self.mlp(self.ln2(x)) class RecordEncoder(nn.Module): """Embed one cache record from its token pieces + typed features. Deliberately shallow: a record is a typed exact value, not a document.""" def __init__(self, d: int, tok_emb: nn.Embedding): super().__init__() self.tok_emb = tok_emb self.store_emb = nn.Embedding(4, d) self.kind_emb = nn.Embedding(8, d) self.key_emb = nn.Embedding(24, d) self.ent_emb = nn.Embedding(65, d) self.age_proj = nn.Linear(8, d) self.mix = nn.Sequential(nn.LayerNorm(d), nn.Linear(d, d), nn.GELU(), nn.Linear(d, d)) self.ln = nn.LayerNorm(d) @staticmethod def age_feats(age: torch.Tensor) -> torch.Tensor: a = age.float().clamp(min=0) cols = [torch.log1p(a)] for s in (8.0, 32.0, 128.0, 512.0): cols += [torch.sin(a / s), torch.cos(a / s)] return torch.stack(cols[:8], dim=-1) def static(self, val_toks, key_toks, store, kind, key_id, ent): """Age-independent part: embed once per chunk/batch.""" # tokens: [.., Tv] / [.., Tk]; averaged piece embeddings (order-light, # values are copied by pointer, never decoded) vmask = (val_toks != 0).float().unsqueeze(-1) kmask = (key_toks != 0).float().unsqueeze(-1) v = (self.tok_emb(val_toks.long()) * vmask).sum(-2) / vmask.sum(-2).clamp(min=1) k = (self.tok_emb(key_toks.long()) * kmask).sum(-2) / kmask.sum(-2).clamp(min=1) h = (v + k + self.store_emb(store.long()) + self.kind_emb(kind.long().clamp(max=7)) + self.key_emb(key_id.long().clamp(max=23)) + self.ent_emb(ent.long().clamp(max=64))) return h + self.mix(h) def finalize(self, h, age): """Add the record's current age (events since birth) and normalize.""" return self.ln(h + self.age_proj(self.age_feats(age))) class HeadSystem(nn.Module): """Factorized outputs with FIXED normalized composition of the three representation paths (state / current event / workspace) - no learned router (hypothesis §5: learned routing collapsed in the parent project).""" GROUPS = ("mode", "enum", "ptr", "op") def __init__(self, d: int): super().__init__() self.P = nn.ModuleDict({f"{g}_{src}": nn.Linear(d, d, bias=False) for g in self.GROUPS for src in ("s", "e", "w")}) self.ln = nn.ModuleDict({g: nn.LayerNorm(d) for g in self.GROUPS}) self.mode = nn.Linear(d, N_MODES) self.enum = nn.Linear(d, N_ENUM) self.ptr_q = nn.Linear(d, d, bias=False) self.ptr_k = nn.Linear(d, d, bias=False) self.op = nn.Linear(d, N_OPS) self.arg_q = nn.Linear(d, 2 * d, bias=False) # two argument pointer queries self.d = d def _h(self, g, r_s, r_e, r_w): z = self.P[f"{g}_e"](r_e) if r_s is not None: z = z + self.P[f"{g}_s"](r_s) if r_w is not None: z = z + self.P[f"{g}_w"](r_w) return self.ln[g](z) def forward(self, r_s, r_e, r_w, rec_bank, live_mask): """rec_bank: [B,S,d] embedded live records (dead rows zero), live_mask: [B,S] bool.""" out = {} out["mode"] = self.mode(self._h("mode", r_s, r_e, r_w)) out["enum"] = self.enum(self._h("enum", r_s, r_e, r_w)) hp = self._h("ptr", r_s, r_e, r_w) keys = self.ptr_k(rec_bank) # [B,S,d] scores = torch.einsum("bd,bsd->bs", self.ptr_q(hp), keys) / math.sqrt(self.d) out["ptr"] = scores.masked_fill(~live_mask, float("-inf")) ho = self._h("op", r_s, r_e, r_w) out["op"] = self.op(ho) aq = self.arg_q(ho).view(-1, 2, self.d) ascore = torch.einsum("bad,bsd->bas", aq, keys) / math.sqrt(self.d) out["args"] = ascore.masked_fill(~live_mask.unsqueeze(1), float("-inf")) return out def enum_legal_mask(bitmask: torch.Tensor) -> torch.Tensor: """[B,8] uint8 -> [B,N_ENUM] bool.""" bits = torch.arange(8, device=bitmask.device) expanded = (bitmask.unsqueeze(-1) >> bits) & 1 # [B,8,8] return expanded.reshape(bitmask.shape[0], 64)[:, :N_ENUM].bool() def losses(out, sup, w_mode=0.3, w_noout=0.05, w_enum=2.0, w_other=1.0): """Shared loss over one flat batch of positions. w_enum/w_other implement the protected-curriculum stage-2 rebalance: the enum (semantic/algorithmic) circuit sits below its learning transition when pointer/op gradients dominate the shared trunk (hypothesis §10).""" device = out["mode"].device mode_gold = sup["mode_gold"].long() weights = torch.where(mode_gold == 0, torch.tensor(w_noout, device=device), torch.tensor(1.0, device=device)) L = {} L["mode"] = (F.cross_entropy(out["mode"], mode_gold, reduction="none") * weights).mean() * w_mode * w_other m_enum = mode_gold == 2 if m_enum.any(): legal = enum_legal_mask(sup["enum_legal"][m_enum]) logits = out["enum"][m_enum].masked_fill(~legal, float("-inf")) L["enum"] = w_enum * F.cross_entropy(logits, sup["enum_gold"][m_enum].long()) m_ptr = mode_gold == 1 if m_ptr.any(): L["ptr"] = w_other * F.cross_entropy(out["ptr"][m_ptr], sup["ptr_gold_slot"][m_ptr].long()) m_op = mode_gold == 3 if m_op.any(): L["op"] = 0.5 * w_other * F.cross_entropy(out["op"][m_op], sup["op_gold"][m_op].long()) for a in range(2): tgt = sup["op_arg_slots"][m_op][:, a].long() va = tgt >= 0 if va.any(): L[f"arg{a}"] = w_other * F.cross_entropy(out["args"][m_op][va][:, a], tgt[va]) return L