pns-bind-25m / src /pns /data /pack.py
nur-dev's picture
PNS-Bind-25M: implementation, configs, eval, results, reproduction
f930dac verified
Raw History Blame Contribute Delete
8.09 kB
"""Encode lifetimes into npz shards.
One shard = many lifetimes, flat arrays + offset indexes. Cache contents are
stored as per-event add/evict deltas against the per-lifetime record table, so
every consumer (PNSR trainer, TX window sampler, nulls, eval) materialises the
exact same bounded stores the generation-time reducer certified.
"""
from __future__ import annotations
import hashlib
import os
import tempfile
from pathlib import Path
import numpy as np
from ..world.reducer import N_PTR_SLOTS
from .binding import N_BIND_SLOTS, binding_plan
from ..world.schema import ET, Event
from .tokenize import encode_event
VAL_T, KEY_T = 12, 14 # token pieces kept per record value/key
MAX_EV_TOKENS = 112
N_ENT_BUCKETS = 64
def val_hash(s: str) -> np.uint64:
return np.uint64(int.from_bytes(hashlib.sha256(s.encode()).digest()[:8], "little"))
def pack_lifetimes(lifetimes: list[dict], tok, out_path: Path) -> dict:
A = {k: [] for k in ("tokens", "etype", "dt", "mode_gold", "family", "enum_gold",
"ptr_gold_slot", "op_gold", "gold_val_hash", "delay", "kappa",
"gold_age_rank", "corrected", "reverted", "chain_len",
"bind_write", "bind_read", "rec_store", "rec_kind",
"rec_ent", "rec_key", "rec_birth_ev", "rec_slot", "rec_val_hash",
"add_rid", "evict_rid")}
ev_tok_off, lt_off, rec_lt_off = [0], [0], [0]
add_off, evict_off = [0], [0]
op_arg_slots, enum_legal = [], []
rec_val_toks, rec_key_toks = [], []
lt_seed, lt_len, lt_digest = [], [], []
n_trunc = 0
slot_ent_all, slot_attr_all = [], []
for lt in lifetimes:
red = lt["reducer"]
slot_hist: dict[int, int] = {} # rid -> stable slot (recorded at add time)
events: list[Event] = lt["events"]
bp = binding_plan(lt)
A["bind_write"].extend(bp["write_slot"].tolist())
A["bind_read"].extend(bp["read_slot"].tolist())
slot_ent_all.append(bp["slot_ent"])
slot_attr_all.append(bp["slot_attr"])
for ev in events:
ids = encode_event(tok, ev)
if len(ids) > MAX_EV_TOKENS:
assert ev.etype != ET.QUESTION, f"question truncated: {ev.text!r}"
ids = ids[:MAX_EV_TOKENS]
n_trunc += 1
A["tokens"].extend(ids)
ev_tok_off.append(len(A["tokens"]))
A["etype"].append(int(ev.etype))
A["dt"].append(ev.dt)
A["mode_gold"].append(int(ev.mode_gold))
A["family"].append(ev.family)
A["enum_gold"].append(ev.enum_gold)
mask = np.zeros(8, np.uint8)
for e in ev.enum_legal:
mask[e >> 3] |= 1 << (e & 7)
enum_legal.append(mask)
# pointer gold -> stable slot id
if ev.ptr_gold_rid >= 0:
A["ptr_gold_slot"].append(slot_hist[ev.ptr_gold_rid])
else:
A["ptr_gold_slot"].append(-1)
A["op_gold"].append(ev.op_gold)
args = [slot_hist[r] for r in ev.op_arg_rids[:2]]
op_arg_slots.append(args + [-1] * (2 - len(args)))
A["gold_val_hash"].append(val_hash(ev.gold_value) if ev.gold_value else np.uint64(0))
m = ev.meta
A["delay"].append(m.get("delay", -1))
A["kappa"].append(m.get("kappa", -1))
A["gold_age_rank"].append(m.get("gold_age_rank", -1))
A["corrected"].append(m.get("corrected", -1))
A["reverted"].append(m.get("reverted", -1))
A["chain_len"].append(m.get("chain_len", -1))
# cache deltas; slots were assigned by the reducer at add time
for rid in m["evicts"]:
A["evict_rid"].append(rid)
for rid, gslot in zip(m["adds"], m["add_slots"]):
slot_hist[rid] = gslot
A["add_rid"].append(rid)
add_off.append(len(A["add_rid"]))
evict_off.append(len(A["evict_rid"]))
lt_off.append(len(A["etype"]))
# record table
for r in red.archive:
A["rec_store"].append(r.store)
A["rec_kind"].append(r.kind)
A["rec_ent"].append(r.ent % N_ENT_BUCKETS)
A["rec_key"].append(r.key_id)
A["rec_birth_ev"].append(r.birth_ev)
A["rec_slot"].append(slot_hist[r.rid])
A["rec_val_hash"].append(val_hash(r.value))
vt = tok.encode(r.value).ids[:VAL_T]
kt = tok.encode(r.key_text).ids[:KEY_T]
rec_val_toks.append(vt + [0] * (VAL_T - len(vt)))
rec_key_toks.append(kt + [0] * (KEY_T - len(kt)))
rec_lt_off.append(len(A["rec_store"]))
lt_seed.append(lt["seed"])
lt_len.append(lt["length"])
lt_digest.append(val_hash(lt["digest"]))
arrays = dict(
tokens=np.asarray(A["tokens"], np.uint16),
ev_tok_off=np.asarray(ev_tok_off, np.int64),
lt_off=np.asarray(lt_off, np.int64),
rec_lt_off=np.asarray(rec_lt_off, np.int64),
add_off=np.asarray(add_off, np.int64),
evict_off=np.asarray(evict_off, np.int64),
etype=np.asarray(A["etype"], np.uint8),
dt=np.asarray(A["dt"], np.uint16),
mode_gold=np.asarray(A["mode_gold"], np.uint8),
family=np.asarray(A["family"], np.int8),
enum_gold=np.asarray(A["enum_gold"], np.int16),
enum_legal=np.stack(enum_legal) if enum_legal else np.zeros((0, 8), np.uint8),
ptr_gold_slot=np.asarray(A["ptr_gold_slot"], np.int16),
op_gold=np.asarray(A["op_gold"], np.int8),
op_arg_slots=np.asarray(op_arg_slots, np.int16),
gold_val_hash=np.asarray(A["gold_val_hash"], np.uint64),
delay=np.asarray(A["delay"], np.int32),
kappa=np.asarray(A["kappa"], np.int16),
gold_age_rank=np.asarray(A["gold_age_rank"], np.int16),
corrected=np.asarray(A["corrected"], np.int8),
bind_write=np.asarray(A["bind_write"], np.int16),
bind_read=np.asarray(A["bind_read"], np.int16),
bind_slot_ent=np.stack(slot_ent_all) if slot_ent_all
else np.zeros((0, N_BIND_SLOTS), np.int16),
bind_slot_attr=np.stack(slot_attr_all) if slot_attr_all
else np.zeros((0, N_BIND_SLOTS), np.int16),
reverted=np.asarray(A["reverted"], np.int8),
chain_len=np.asarray(A["chain_len"], np.int16),
rec_store=np.asarray(A["rec_store"], np.uint8),
rec_kind=np.asarray(A["rec_kind"], np.uint8),
rec_ent=np.asarray(A["rec_ent"], np.int16),
rec_key=np.asarray(A["rec_key"], np.uint8),
rec_birth_ev=np.asarray(A["rec_birth_ev"], np.int32),
rec_slot=np.asarray(A["rec_slot"], np.int16),
rec_val_hash=np.asarray(A["rec_val_hash"], np.uint64),
rec_val_toks=np.asarray(rec_val_toks, np.uint16) if rec_val_toks
else np.zeros((0, VAL_T), np.uint16),
rec_key_toks=np.asarray(rec_key_toks, np.uint16) if rec_key_toks
else np.zeros((0, KEY_T), np.uint16),
add_rid=np.asarray(A["add_rid"], np.int32),
evict_rid=np.asarray(A["evict_rid"], np.int32),
lt_seed=np.asarray(lt_seed, np.int64),
lt_len=np.asarray(lt_len, np.int32),
lt_digest=np.asarray(lt_digest, np.uint64),
)
assert arrays["ev_tok_off"][-1] == len(arrays["tokens"])
out_path.parent.mkdir(parents=True, exist_ok=True)
fd, tmp = tempfile.mkstemp(dir=out_path.parent, suffix=".tmp")
os.close(fd)
np.savez(tmp, **arrays) # uncompressed: memmap-friendly at load
saved = tmp if tmp.endswith(".npz") else tmp + ".npz" # numpy appends .npz
os.replace(saved, out_path)
if os.path.exists(tmp):
os.unlink(tmp)
return dict(n_lifetimes=len(lifetimes), n_events=int(arrays["etype"].shape[0]),
n_tokens=int(arrays["tokens"].shape[0]),
n_records=int(arrays["rec_store"].shape[0]), n_trunc=n_trunc,
sha=hashlib.sha256(out_path.read_bytes()).hexdigest()[:16])
assert N_PTR_SLOTS == 100