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