Download scripts/audit_star_data.py from Avra98/latent_backtrack: direct link, hf CLI and curl.
- Browser
- Download file 3.49 kB
-
https://huggingface.co/Avra98/latent_backtrack/resolve/main/scripts/audit_star_data.py
- Command line
-
hf download hf://Avra98/latent_backtrack/scripts/audit_star_data.py
-
curl -L -o audit_star_data.py https://huggingface.co/Avra98/latent_backtrack/resolve/main/scripts/audit_star_data.py
3.49 kB
| #!/usr/bin/env python3 | |
| """Audit 2-arm star datasets for structural bugs / leakage / distribution shifts. | |
| For each depth: validate every train+valid sample's graph structure, check | |
| train<->valid leakage at the graph level AND the hop-1-pattern level, and | |
| report stats that could differ between working (L6/L10) and failing (L12/L14) | |
| depths. | |
| """ | |
| import json | |
| import sys | |
| from collections import Counter | |
| def load(path): | |
| return json.load(open(path)) | |
| def audit_depth(L): | |
| print(f"\n================ L{L} ================") | |
| tr = load(f"data/star_2arm_L{L}_train_fo_bfs.json") | |
| va = load(f"data/star_2arm_L{L}_valid_fo_bfs.json") | |
| n_nodes = 2 + 4 * L | |
| bad = Counter() | |
| def check(s): | |
| edges = [tuple(e) for e in s["edges"]] | |
| if len(edges) != 4 * L: | |
| bad["edge_count"] += 1 | |
| root, neg_root = s["root"], s["neg_root"] | |
| # adjacency (directed src->dst as generated) | |
| out = {} | |
| for a, b in edges: | |
| out.setdefault(a, []).append(b) | |
| # root must have exactly 2 outgoing edges (2 arms) | |
| if sorted(out.get(root, [])) != sorted(s["neighbor_k"]["1"]): | |
| bad["hop1_mismatch"] += 1 | |
| if len(out.get(root, [])) != 2: | |
| bad["root_degree"] += 1 | |
| # BFS frontier sets: neighbor_k[k] must be reachable at exactly k hops | |
| frontier = {root} | |
| for k in range(1, L + 1): | |
| nxt = set() | |
| for v in frontier: | |
| nxt.update(out.get(v, [])) | |
| if set(s["neighbor_k"][str(k)]) != nxt: | |
| bad[f"frontier_k"] += 1 | |
| break | |
| frontier = nxt | |
| # target must be a leaf at hop L; neg_target unreachable from root | |
| reach = {root} | |
| stack = [root] | |
| while stack: | |
| v = stack.pop() | |
| for w in out.get(v, []): | |
| if w not in reach: | |
| reach.add(w) | |
| stack.append(w) | |
| if s["target"] not in reach: | |
| bad["target_unreachable"] += 1 | |
| if s["neg_target"] in reach: | |
| bad["neg_target_reachable"] += 1 | |
| if len(reach) != 1 + 2 * L: | |
| bad["component_size"] += 1 | |
| for s in tr + va: | |
| check(s) | |
| # graph-level leakage | |
| def key(s): | |
| return (s["root"], s["target"], s["neg_target"], | |
| frozenset(tuple(e) for e in s["edges"])) | |
| tr_keys = {key(s) for s in tr} | |
| va_keys = {key(s) for s in va} | |
| leak = len(tr_keys & va_keys) | |
| # hop-1 pattern coverage: (root, frozenset(hop1 neighbors)) | |
| tr_h1 = {(s["root"], frozenset(s["neighbor_k"]["1"])) for s in tr} | |
| va_h1 = [(s["root"], frozenset(s["neighbor_k"]["1"])) for s in va] | |
| seen_h1 = sum(1 for h in va_h1 if h in tr_h1) | |
| ids = Counter() | |
| for s in tr: | |
| for x in s["idx_to_symbol"]: | |
| ids[int(x)] += 1 | |
| pool = max(ids) + 1 | |
| print(f"train={len(tr)} valid={len(va)} nodes/sample={n_nodes} pool={pool}") | |
| print(f"structural violations: {dict(bad) if bad else 'NONE'}") | |
| print(f"graph-level train/val overlap: {leak}") | |
| print(f"unique hop-1 patterns in train: {len(tr_h1)}") | |
| print(f"val hop-1 patterns also present in train: {seen_h1}/{len(va)} " | |
| f"({seen_h1/len(va):.1%})") | |
| mn, mx = min(ids.values()), max(ids.values()) | |
| print(f"id usage min/max across pool: {mn}/{mx} (ratio {mn/mx:.2f})") | |
| if __name__ == "__main__": | |
| for L in (int(x) for x in (sys.argv[1:] or ["6", "10", "12", "14"])): | |
| audit_depth(L) | |