pns-bind-25m / src /pns /data /view.py
nur-dev's picture
PNS-Bind-25M: implementation, configs, eval, results, reproduction
f930dac verified
Raw History Blame Contribute Delete
4.6 kB
"""Read-side view of packed shards: per-lifetime iteration with incremental
live-cache materialisation. Single implementation shared by the null attacks,
the TX window sampler, the PNSR streaming batcher and the evaluators, so every
consumer sees byte-identical cache states.
"""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
import numpy as np
from ..world.reducer import N_PTR_SLOTS
# Split identifier -> directory inside the published dataset repo. The
# identifiers are the ones used throughout the study and in every result file;
# the directory names are the published layout.
SPLIT_DIRS = {
"train": "e1_original/train",
"val": "e1_original/val",
"val_1k": "e1_original/val_1k",
"val_4k": "e1_original/val_4k",
"test": "e1_original/test",
"test_1k": "e1_original/test_1k",
"test_4k": "e1_original/test_4k",
"train_relabel": "e1_corrected_eval/train_relabel",
"val_relabel": "e1_corrected_eval/val_relabel",
"e3_train": "e3_corrected/train",
"e3_dev": "e3_corrected/dev",
"e3_dev_1k": "e3_corrected/dev_1k",
"e3_conf": "e3_corrected/confirmation",
"e3_conf_1k": "e3_corrected/confirmation_1k",
}
def shard_paths(split: str, root: Path) -> list[Path]:
"""Shards of `split` under `root` (= <dataset repo>/data).
Accepts either a study split identifier (see SPLIT_DIRS) or a literal
subdirectory name, so a locally regenerated corpus also works.
"""
d = root / SPLIT_DIRS.get(split, split)
if not d.is_dir():
d = root / split
return sorted(d.glob("shard_*.npz"))
class Shard:
def __init__(self, path: Path):
self.z = np.load(path, mmap_mode=None) # npz members load lazily on access
self.path = path
for k in self.z.files:
setattr(self, k, self.z[k])
self.n_lifetimes = len(self.lt_off) - 1
def lifetime(self, i: int) -> "LifetimeView":
return LifetimeView(self, i)
@dataclass
class EventView:
idx: int # event index within lifetime
tokens: np.ndarray
etype: int
dt: int
live_slots: np.ndarray # int32 [N_PTR_SLOTS] -> record row (global) or -1
mode_gold: int
family: int
enum_gold: int
enum_legal: np.ndarray # uint8 [8] bitmask
ptr_gold_slot: int
op_gold: int
op_arg_slots: np.ndarray
meta: dict
class LifetimeView:
def __init__(self, sh: Shard, i: int):
self.sh = sh
self.lo = int(sh.lt_off[i])
self.hi = int(sh.lt_off[i + 1])
self.rec_lo = int(sh.rec_lt_off[i])
self.rec_hi = int(sh.rec_lt_off[i + 1])
self.seed = int(sh.lt_seed[i])
self.length = self.hi - self.lo
def records(self) -> dict[str, np.ndarray]:
s = slice(self.rec_lo, self.rec_hi)
sh = self.sh
return dict(store=sh.rec_store[s], kind=sh.rec_kind[s], ent=sh.rec_ent[s],
key=sh.rec_key[s], birth_ev=sh.rec_birth_ev[s], slot=sh.rec_slot[s],
val_hash=sh.rec_val_hash[s], val_toks=sh.rec_val_toks[s],
key_toks=sh.rec_key_toks[s])
def __iter__(self):
sh = self.sh
live = np.full(N_PTR_SLOTS, -1, np.int32)
for e in range(self.lo, self.hi):
for rid in sh.evict_rid[sh.evict_off[e]:sh.evict_off[e + 1]]:
slot = sh.rec_slot[self.rec_lo + int(rid)]
live[slot] = -1
for rid in sh.add_rid[sh.add_off[e]:sh.add_off[e + 1]]:
slot = sh.rec_slot[self.rec_lo + int(rid)]
live[slot] = self.rec_lo + int(rid) # global record row
yield EventView(
idx=e - self.lo,
tokens=sh.tokens[sh.ev_tok_off[e]:sh.ev_tok_off[e + 1]],
etype=int(sh.etype[e]), dt=int(sh.dt[e]),
live_slots=live.copy(),
mode_gold=int(sh.mode_gold[e]), family=int(sh.family[e]),
enum_gold=int(sh.enum_gold[e]), enum_legal=sh.enum_legal[e],
ptr_gold_slot=int(sh.ptr_gold_slot[e]), op_gold=int(sh.op_gold[e]),
op_arg_slots=sh.op_arg_slots[e],
meta=dict(delay=int(sh.delay[e]), kappa=int(sh.kappa[e]),
gold_age_rank=int(sh.gold_age_rank[e]),
corrected=int(sh.corrected[e]), reverted=int(sh.reverted[e]),
chain_len=int(sh.chain_len[e]),
gold_val_hash=int(sh.gold_val_hash[e])),
)