"""PNS-Bind (E3): exact identity addressing, learned semantic content. Seven prior models kept a learned recurrent state that was causally inert. E3 tests one narrow hypothesis: learned semantic memory becomes durable when exact identity/binding structure protects it from unrelated updates. mode="bind" : an event writes ONLY the slot its exact address names; every other slot is carried through BITWISE unchanged. mode="unbound" : identical capacity (192 slots), identical everything else, but the original unstructured mechanism - every slot is written every event by attention. This is the E3-A3 capacity-matched control that separates "binding isolation" from "more learned-state capacity". The exact side supplies only identity: slot -> (entity id, attribute type) and "which slot is authoritative for this query". It never supplies a value; the payload s_i in R^d is learned, and the answer must be decoded from it. """ from __future__ import annotations from dataclasses import dataclass import torch import torch.nn as nn from ..data.binding import N_BIND_SLOTS from ..world.schema import ET, ENUM_VOCAB from .modules import SDPA, Block, HeadSystem, RecordEncoder, add_mask @dataclass class BindConfig: vocab: int = 8192 d: int = 192 heads: int = 6 enc_layers: int = 3 n_slots: int = N_BIND_SLOTS T_ev: int = 48 mode: str = "bind" # "bind" | "unbound" gate_bias: float = 0.0 full_heads: bool = False # Stage B: the complete factorized output system class PNSBind(nn.Module): def __init__(self, cfg: BindConfig): super().__init__() self.cfg = cfg d = cfg.d self.tok_emb = nn.Embedding(cfg.vocab, d) self.pos_emb = nn.Parameter(torch.zeros(cfg.T_ev, d)) self.enc = nn.ModuleList(Block(d, cfg.heads) for _ in range(cfg.enc_layers)) self.enc_ln = nn.LayerNorm(d) self.etype_emb = nn.Embedding(len(ET), d) self.dt_proj = nn.Linear(8, d) self.slot0 = nn.Parameter(torch.randn(cfg.n_slots, d) * 0.02) # exact identity features of a slot: entity and attribute type only self.key_ent = nn.Embedding(4097, d) self.key_attr = nn.Embedding(32, d) self.wq = nn.Linear(d, d, bias=False) self.write_attn = SDPA(d, cfg.heads) self.gate = nn.Sequential(nn.Linear(3 * d, d), nn.GELU(), nn.Linear(d, d)) nn.init.constant_(self.gate[-1].bias, cfg.gate_bias) self.read_q = nn.Parameter(torch.randn(1, d) * 0.02) self.hop = SDPA(d, cfg.heads) # learned hop 2 over slot keys self.mix = nn.Sequential(nn.LayerNorm(3 * d), nn.Linear(3 * d, d), nn.GELU(), nn.Linear(d, d)) self.enum = nn.Linear(d, len(ENUM_VOCAB)) if cfg.full_heads: # identical head system to every prior model, so Stage B differs # from the failed runs ONLY in the state mechanism self.recenc = RecordEncoder(d, self.tok_emb) self.heads = HeadSystem(d) def initial_state(self, B, device): return self.slot0.detach().float().unsqueeze(0).repeat(B, 1, 1).to(device) def encode_event(self, tok, etype, dt): pad = tok == 0 x = self.tok_emb(tok.long()) + self.pos_emb[: tok.shape[1]] for blk in self.enc: x = blk(x, self_padding_mask=pad) x = self.enc_ln(x) den = (~pad).sum(1, keepdim=True).clamp(min=1) v = x.masked_fill(pad.unsqueeze(-1), 0).sum(1) / den v = v + self.etype_emb(etype.long()) + self.dt_proj(RecordEncoder.age_feats(dt)) return x, v, pad def step(self, state, tok, etype, dt, bind_write, bind_read, slot_ent, slot_attr, freeze_writes=False, payload_scale=1.0, rec_bank=None, live_mask=None): """state [B,S,d] fp32. bind_write/bind_read: [B] int, -1 = none.""" cfg = self.cfg B, S, d = state.shape x, v, pad = self.encode_event(tok, etype, dt) cdt = x.dtype keys = self.key_ent((slot_ent.long() + 1).clamp(0, 4096)) \ + self.key_attr(slot_attr.long().clamp(0, 31)) if not freeze_writes: st = state.to(cdt) if cfg.mode == "bind": # ---- addressed write: gather ONE slot, update it, scatter back. # Slots other than bind_write are never touched, so they are # carried through bitwise (asserted in tests). w = bind_write.long() has = (w >= 0) idx = w.clamp(min=0).view(B, 1, 1).expand(-1, 1, d) cur = torch.gather(st, 1, idx) # [B,1,d] kcur = torch.gather(keys.to(cdt), 1, idx) q = self.wq(cur + kcur) kv = torch.cat([x, v.unsqueeze(1)], 1) kvm = torch.cat([pad, torch.zeros(B, 1, dtype=torch.bool, device=x.device)], 1) delta = self.write_attn(q, kv, kv, mask=add_mask(kvm, None, q.dtype)) g = torch.sigmoid(self.gate(torch.cat([cur, delta, v.unsqueeze(1)], -1))) new = ((1 - g) * cur + g * delta).float() upd = state.scatter(1, w.clamp(min=0).view(B, 1, 1).expand(-1, 1, d), new) state = torch.where(has.view(B, 1, 1), upd, state) else: # ---- unstructured control: every slot written every event q = self.wq(st + keys.to(cdt)) kv = torch.cat([x, v.unsqueeze(1)], 1) kvm = torch.cat([pad, torch.zeros(B, 1, dtype=torch.bool, device=x.device)], 1) delta = self.write_attn(q, kv, kv, mask=add_mask(kvm, None, q.dtype)) g = torch.sigmoid(self.gate(torch.cat( [st, delta, v.unsqueeze(1).expand(-1, S, -1)], -1))) state = ((1 - g) * st + g * delta).float() # ---------------- read st = (state * payload_scale).to(cdt) kslots = st + keys.to(cdt) r = bind_read.long() addressed = torch.gather( st, 1, r.clamp(min=0).view(B, 1, 1).expand(-1, 1, d)) addressed = addressed * (r >= 0).view(B, 1, 1).to(cdt) # [B,1,d] # learned second hop: query the slot bank with what hop 1 produced hop2 = self.hop(addressed + v.unsqueeze(1), kslots, kslots) dense = (st * 0 + kslots).mean(1, keepdim=True) # dense read h = self.mix(torch.cat([addressed.squeeze(1), hop2.squeeze(1), v], -1)) if self.cfg.full_heads and rec_bank is not None: out = self.heads(h, v, kslots.mean(1), rec_bank, live_mask) out["enum"] = out["enum"] + self.enum(h) # binding-read pathway return state, out return state, {"enum": self.enum(h)}