pns-bind-25m / src /pns /model /bind.py
nur-dev's picture
PNS-Bind-25M: implementation, configs, eval, results, reproduction
f930dac verified
Raw History Blame Contribute Delete
7.12 kB
"""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)}