"""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` (= /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])), )