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