Download src/pns/model/bind.py from nur-dev/pns-bind-25m: direct link, hf CLI and curl.
- Browser
- Download file 7.12 kB
-
https://huggingface.co/nur-dev/pns-bind-25m/resolve/main/src/pns/model/bind.py
- Command line
-
hf download hf://nur-dev/pns-bind-25m/src/pns/model/bind.py
-
curl -L -o bind.py https://huggingface.co/nur-dev/pns-bind-25m/resolve/main/src/pns/model/bind.py
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 | |
| 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)} | |