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