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