audit: vectorised rolling-hash grams + sorted-array lookup; the Python-set version could not finish 1.2B tokens
Browse files- build/audit_contamination.py +135 -82
build/audit_contamination.py
CHANGED
|
@@ -6,34 +6,40 @@
|
|
| 6 |
# hard: the only values that ever leave the process are counts and rates per source. No item text, no
|
| 7 |
# lengths, no samples, nothing a human could read as a question. The print surface is a dict of integers.
|
| 8 |
#
|
| 9 |
-
# Which splits are
|
| 10 |
-
#
|
| 11 |
-
#
|
| 12 |
-
#
|
| 13 |
-
# {train, validation, dev} where present, per task, and NEVER `test`. That is more conservative than the
|
| 14 |
-
# literal wording, and it is the set that actually matters.
|
| 15 |
#
|
| 16 |
-
#
|
| 17 |
-
#
|
| 18 |
-
#
|
| 19 |
-
#
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 20 |
#
|
| 21 |
-
#
|
| 22 |
-
# held as a Python set of ints. The mix side streams document by document from the shard files, so its
|
| 23 |
-
# size is irrelevant. Nothing here holds a billion of anything.
|
| 24 |
-
#
|
| 25 |
-
# Read-only over the mix. CPU session. No credentials.
|
| 26 |
|
| 27 |
import argparse
|
| 28 |
-
import hashlib
|
| 29 |
import json
|
| 30 |
import os
|
| 31 |
import struct
|
| 32 |
|
| 33 |
-
import
|
| 34 |
|
| 35 |
K = 13
|
| 36 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 37 |
REFS = [
|
| 38 |
("arc_easy", "allenai/ai2_arc", ["ARC-Easy"], ["train", "validation"]),
|
| 39 |
("arc_challenge", "allenai/ai2_arc", ["ARC-Challenge"], ["train", "validation"]),
|
|
@@ -47,38 +53,56 @@ REFS = [
|
|
| 47 |
TEXT_KEYS = ("question", "text", "choices", "story", "sentence", "prompt", "ctx", "input", "query")
|
| 48 |
|
| 49 |
|
| 50 |
-
def
|
| 51 |
-
"""
|
| 52 |
-
|
| 53 |
-
|
| 54 |
-
|
| 55 |
-
|
| 56 |
-
|
| 57 |
-
|
| 58 |
-
|
| 59 |
-
|
| 60 |
-
|
| 61 |
-
|
| 62 |
-
|
| 63 |
-
|
| 64 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 65 |
if isinstance(v, str):
|
| 66 |
-
|
| 67 |
elif isinstance(v, dict):
|
| 68 |
for x in v.values():
|
| 69 |
-
flatten(x,
|
| 70 |
elif isinstance(v, (list, tuple)):
|
| 71 |
for x in v:
|
| 72 |
-
flatten(x,
|
| 73 |
|
| 74 |
|
| 75 |
def build_reference(tk, max_items_per_split):
|
| 76 |
-
"""
|
| 77 |
-
reference
|
| 78 |
import datasets
|
| 79 |
|
| 80 |
-
|
| 81 |
-
|
|
|
|
|
|
|
| 82 |
for cfg in configs:
|
| 83 |
for sp in splits:
|
| 84 |
key = f"{task}::{cfg or 'default'}::{sp}"
|
|
@@ -86,10 +110,10 @@ def build_reference(tk, max_items_per_split):
|
|
| 86 |
ds = (datasets.load_dataset(repo, cfg, split=sp, streaming=True) if cfg
|
| 87 |
else datasets.load_dataset(repo, split=sp, streaming=True))
|
| 88 |
except Exception as e:
|
| 89 |
-
|
| 90 |
print(f" ref {key}: UNREADABLE {type(e).__name__}", flush=True)
|
| 91 |
continue
|
| 92 |
-
g, n =
|
| 93 |
for row in ds:
|
| 94 |
if n >= max_items_per_split:
|
| 95 |
break
|
|
@@ -101,16 +125,30 @@ def build_reference(tk, max_items_per_split):
|
|
| 101 |
if not txt.strip():
|
| 102 |
continue
|
| 103 |
n += 1
|
| 104 |
-
|
| 105 |
-
|
| 106 |
-
|
| 107 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 108 |
|
| 109 |
|
| 110 |
def main():
|
| 111 |
ap = argparse.ArgumentParser()
|
| 112 |
ap.add_argument("--root", default="/kaggle/working/mixroot")
|
| 113 |
ap.add_argument("--max-items-per-split", type=int, default=20000)
|
|
|
|
|
|
|
| 114 |
a = ap.parse_args()
|
| 115 |
|
| 116 |
from tokenizers import Tokenizer
|
|
@@ -118,70 +156,85 @@ def main():
|
|
| 118 |
tk = Tokenizer.from_file(hf_hub_download("HuggingFaceTB/SmolLM2-135M", "tokenizer.json"))
|
| 119 |
tk.no_truncation()
|
| 120 |
|
| 121 |
-
print("building reference gram
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
| 125 |
|
| 126 |
-
# stream the mix's own shards; per-shard source lists let a hit be attributed to a source, which is
|
| 127 |
-
# the only actionable granularity ("drop or clean any source that fails")
|
| 128 |
man = json.load(open(os.path.join(a.root, "manifest.json")))
|
| 129 |
-
|
|
|
|
|
|
|
| 130 |
hits_by_source = {}
|
| 131 |
docs = grams_seen = 0
|
| 132 |
missing = []
|
| 133 |
-
import build_mix as BM
|
| 134 |
for rec in man["shards"]:
|
| 135 |
p = os.path.join(a.root, rec["file"])
|
| 136 |
if not os.path.exists(p):
|
| 137 |
-
# A shard
|
| 138 |
-
# audit over an empty directory report zero contamination (E-017's shape).
|
| 139 |
missing.append(rec["file"])
|
| 140 |
continue
|
| 141 |
for d in BM.iter_shard_docs(p):
|
| 142 |
docs += 1
|
| 143 |
-
|
| 144 |
-
|
|
|
|
| 145 |
continue
|
| 146 |
-
grams_seen +=
|
| 147 |
-
|
| 148 |
-
|
| 149 |
-
|
| 150 |
-
|
| 151 |
-
|
| 152 |
-
|
| 153 |
-
|
| 154 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 155 |
print(f" audited {docs:,} docs, {grams_seen:,} mix grams", flush=True)
|
| 156 |
|
| 157 |
total_overlap = sum(hits.values())
|
|
|
|
| 158 |
report = {
|
| 159 |
"k": K,
|
| 160 |
-
"
|
|
|
|
| 161 |
"reference_unreadable": errors,
|
|
|
|
|
|
|
| 162 |
"shards_expected": len(man["shards"]),
|
| 163 |
"shards_missing": missing,
|
| 164 |
"mix_documents_audited": docs,
|
| 165 |
"mix_grams_audited": grams_seen,
|
| 166 |
-
"
|
| 167 |
-
"overlap_rate_by_ref": {k: round(v / max(1, grams_seen), 8) for k, v in hits.items()},
|
| 168 |
"overlap_by_source": hits_by_source,
|
| 169 |
-
|
| 170 |
-
|
| 171 |
-
"
|
| 172 |
-
"
|
| 173 |
-
|
| 174 |
-
|
| 175 |
-
|
| 176 |
-
"
|
| 177 |
-
|
| 178 |
-
|
|
|
|
| 179 |
}
|
|
|
|
| 180 |
print(f"VERDICT AUDIT_PASSED={report['AUDIT_PASSED']} docs={docs:,} overlap={total_overlap} "
|
| 181 |
-
f"
|
|
|
|
| 182 |
print("AUDIT_JSON_BEGIN")
|
| 183 |
-
print(
|
| 184 |
print("AUDIT_JSON_END")
|
|
|
|
|
|
|
|
|
|
| 185 |
raise SystemExit(0 if report["AUDIT_PASSED"] else (6 if report["MEASURED"] else 7))
|
| 186 |
|
| 187 |
|
|
|
|
| 6 |
# hard: the only values that ever leave the process are counts and rates per source. No item text, no
|
| 7 |
# lengths, no samples, nothing a human could read as a question. The print surface is a dict of integers.
|
| 8 |
#
|
| 9 |
+
# Which splits are the reference, and why it is broader than the minimum: for most of these tasks the
|
| 10 |
+
# harness draws its few-shot exemplars from TRAIN (ARC, HellaSwag, PIQA, Winogrande) or from MMLU's `dev`,
|
| 11 |
+
# so overlap with those is contamination by construction. The reference set is {train, validation, dev}
|
| 12 |
+
# where present, per task, and NEVER `test`.
|
|
|
|
|
|
|
| 13 |
#
|
| 14 |
+
# Why this file does not use Python sets of hashed grams (it did, and it was wrong): 1.23 B tokens in
|
| 15 |
+
# 13-token windows is ~1.2e9 windows, and `hashlib.sha256` per window in CPython is ~1-2 us each, so the
|
| 16 |
+
# audit would have taken hours to days and simply never finished inside a CPU session. Two changes make it
|
| 17 |
+
# tractable without weakening it:
|
| 18 |
+
# 1. Gram fingerprints are a vectorised multiplicative rolling hash (mod 2^64) computed per document with
|
| 19 |
+
# numpy, not a Python loop. Cost is ~20 uint64 ops per token, in C.
|
| 20 |
+
# 2. The reference is ONE sorted uint64 array plus a uint16 bitmask giving the tasks each gram belongs to,
|
| 21 |
+
# queried with np.searchsorted. Memory is ~1 GB, not the ~7 GB a Python set of 90M ints needs, and
|
| 22 |
+
# per-reference sets are unnecessary.
|
| 23 |
+
# Collision behaviour is stated rather than hidden: 64-bit fingerprints over ~1.2e9 mix grams and ~1e8
|
| 24 |
+
# reference grams give an expected spurious-match count of mix_grams x ref_grams / 2^64, reported as
|
| 25 |
+
# `expected_false_positives`. A raw count below that is noise, not evidence of contamination; above it,
|
| 26 |
+
# the offending source is investigated.
|
| 27 |
#
|
| 28 |
+
# Read-only over the mix. CPU session. No credentials in this file.
|
|
|
|
|
|
|
|
|
|
|
|
|
| 29 |
|
| 30 |
import argparse
|
|
|
|
| 31 |
import json
|
| 32 |
import os
|
| 33 |
import struct
|
| 34 |
|
| 35 |
+
import numpy as np
|
| 36 |
|
| 37 |
K = 13
|
| 38 |
+
P = np.uint64(0x100000001B3) # FNV-1a prime, odd, so invertible mod 2^64
|
| 39 |
+
PINV = np.uint64(pow(int(P), -1, 1 << 64))
|
| 40 |
+
MODMASK = np.uint64((1 << 64) - 1)
|
| 41 |
+
|
| 42 |
+
# task -> (repo, config(s), splits treated as reference material; 'test' deliberately absent)
|
| 43 |
REFS = [
|
| 44 |
("arc_easy", "allenai/ai2_arc", ["ARC-Easy"], ["train", "validation"]),
|
| 45 |
("arc_challenge", "allenai/ai2_arc", ["ARC-Challenge"], ["train", "validation"]),
|
|
|
|
| 53 |
TEXT_KEYS = ("question", "text", "choices", "story", "sentence", "prompt", "ctx", "input", "query")
|
| 54 |
|
| 55 |
|
| 56 |
+
def window_hashes(ids, k=K):
|
| 57 |
+
"""uint64 fingerprint of every k-length window of `ids` (uint16 array), vectorised.
|
| 58 |
+
|
| 59 |
+
h_i = sum_{j=i}^{i+k-1} x_j * P^(i+k-1-j) mod 2^64. Via the prefix recurrence
|
| 60 |
+
C_i = C_{i-1}*P + x_i and the identity h_i = C_{i+k-1} - C_{i-1} * P^k. C is computed as
|
| 61 |
+
P^i * cumsum(x_j * PINV^j), which is what makes it a few numpy passes instead of a Python loop.
|
| 62 |
+
"""
|
| 63 |
+
n = int(ids.shape[0])
|
| 64 |
+
if n < k:
|
| 65 |
+
return np.empty(0, dtype=np.uint64)
|
| 66 |
+
x = ids.astype(np.uint64)
|
| 67 |
+
|
| 68 |
+
def powers(base, m):
|
| 69 |
+
# [base^0, base^1, ..., base^(m-1)] mod 2^64. np.cumprod of a constant vector would start at
|
| 70 |
+
# base^1, which shifts every window hash and silently breaks the match against the reference.
|
| 71 |
+
out = np.empty(m, dtype=np.uint64)
|
| 72 |
+
out[0] = np.uint64(1)
|
| 73 |
+
if m > 1:
|
| 74 |
+
out[1:] = np.cumprod(np.full(m - 1, base, dtype=np.uint64))
|
| 75 |
+
return out
|
| 76 |
+
|
| 77 |
+
pow_p, pow_pi = powers(P, n), powers(PINV, n)
|
| 78 |
+
s = np.cumsum(x * pow_pi, dtype=np.uint64)
|
| 79 |
+
c = s * pow_p # c[i] = C_i, prefix hash ending at i
|
| 80 |
+
pk = np.uint64(pow(int(P), k, 1 << 64))
|
| 81 |
+
prev = np.concatenate((np.zeros(1, dtype=np.uint64), (c[:n - k] * pk)))
|
| 82 |
+
h = (c[k - 1:] - prev[:n - k + 1]) & MODMASK
|
| 83 |
+
return h
|
| 84 |
+
|
| 85 |
+
|
| 86 |
+
def flatten(v, out):
|
| 87 |
if isinstance(v, str):
|
| 88 |
+
out.append(v)
|
| 89 |
elif isinstance(v, dict):
|
| 90 |
for x in v.values():
|
| 91 |
+
flatten(x, out)
|
| 92 |
elif isinstance(v, (list, tuple)):
|
| 93 |
for x in v:
|
| 94 |
+
flatten(x, out)
|
| 95 |
|
| 96 |
|
| 97 |
def build_reference(tk, max_items_per_split):
|
| 98 |
+
"""Returns (sorted uint64 array of every reference gram, parallel uint16 task-bitmask, per-task
|
| 99 |
+
item counts, per-reference errors). Only counts and hashes are ever held or printed."""
|
| 100 |
import datasets
|
| 101 |
|
| 102 |
+
grams, masks = [], []
|
| 103 |
+
counts, errors = {}, {}
|
| 104 |
+
for ti, (task, repo, configs, splits) in enumerate(REFS):
|
| 105 |
+
bit = np.uint16(1 << ti)
|
| 106 |
for cfg in configs:
|
| 107 |
for sp in splits:
|
| 108 |
key = f"{task}::{cfg or 'default'}::{sp}"
|
|
|
|
| 110 |
ds = (datasets.load_dataset(repo, cfg, split=sp, streaming=True) if cfg
|
| 111 |
else datasets.load_dataset(repo, split=sp, streaming=True))
|
| 112 |
except Exception as e:
|
| 113 |
+
errors[key] = f"{type(e).__name__}: {str(e)[:140]}"
|
| 114 |
print(f" ref {key}: UNREADABLE {type(e).__name__}", flush=True)
|
| 115 |
continue
|
| 116 |
+
g, n = [], 0
|
| 117 |
for row in ds:
|
| 118 |
if n >= max_items_per_split:
|
| 119 |
break
|
|
|
|
| 125 |
if not txt.strip():
|
| 126 |
continue
|
| 127 |
n += 1
|
| 128 |
+
ids = np.asarray(tk.encode(txt, add_special_tokens=False).ids, dtype=np.uint16)
|
| 129 |
+
h = window_hashes(ids)
|
| 130 |
+
if h.size:
|
| 131 |
+
g.append(h)
|
| 132 |
+
if g:
|
| 133 |
+
arr = np.concatenate(g)
|
| 134 |
+
grams.append(arr)
|
| 135 |
+
masks.append(np.full(arr.shape, bit, dtype=np.uint16))
|
| 136 |
+
counts[key] = n
|
| 137 |
+
print(f" ref {key}: {n} items -> {sum(a.size for a in g):,} grams", flush=True)
|
| 138 |
+
if not grams:
|
| 139 |
+
raise SystemExit("no reference grams built at all -- the audit would be a pass over nothing")
|
| 140 |
+
allg = np.concatenate(grams)
|
| 141 |
+
allm = np.concatenate(masks)
|
| 142 |
+
order = np.argsort(allg, kind="stable")
|
| 143 |
+
return allg[order], allm[order], counts, errors
|
| 144 |
|
| 145 |
|
| 146 |
def main():
|
| 147 |
ap = argparse.ArgumentParser()
|
| 148 |
ap.add_argument("--root", default="/kaggle/working/mixroot")
|
| 149 |
ap.add_argument("--max-items-per-split", type=int, default=20000)
|
| 150 |
+
ap.add_argument("--out", default="", help="write the report JSON here (publish_mix expects "
|
| 151 |
+
"mixroot/audit.json)")
|
| 152 |
a = ap.parse_args()
|
| 153 |
|
| 154 |
from tokenizers import Tokenizer
|
|
|
|
| 156 |
tk = Tokenizer.from_file(hf_hub_download("HuggingFaceTB/SmolLM2-135M", "tokenizer.json"))
|
| 157 |
tk.no_truncation()
|
| 158 |
|
| 159 |
+
print("building reference gram table (only counts ever leave this process)", flush=True)
|
| 160 |
+
ref, refmask, refcounts, errors = build_reference(tk, a.max_items_per_split)
|
| 161 |
+
print(f"reference: {ref.size:,} grams, {len(set(refcounts))} splits ok, {len(errors)} unreadable",
|
| 162 |
+
flush=True)
|
| 163 |
|
|
|
|
|
|
|
| 164 |
man = json.load(open(os.path.join(a.root, "manifest.json")))
|
| 165 |
+
import build_mix as BM
|
| 166 |
+
|
| 167 |
+
hits = {}
|
| 168 |
hits_by_source = {}
|
| 169 |
docs = grams_seen = 0
|
| 170 |
missing = []
|
|
|
|
| 171 |
for rec in man["shards"]:
|
| 172 |
p = os.path.join(a.root, rec["file"])
|
| 173 |
if not os.path.exists(p):
|
| 174 |
+
# A shard not on disk is unmeasured, not clean (E-017's shape).
|
|
|
|
| 175 |
missing.append(rec["file"])
|
| 176 |
continue
|
| 177 |
for d in BM.iter_shard_docs(p):
|
| 178 |
docs += 1
|
| 179 |
+
ids = np.frombuffer(bytes(d), dtype=np.uint16)
|
| 180 |
+
g = window_hashes(ids)
|
| 181 |
+
if g.size == 0:
|
| 182 |
continue
|
| 183 |
+
grams_seen += g.size
|
| 184 |
+
pos = np.searchsorted(ref, g)
|
| 185 |
+
pos_c = np.clip(pos, 0, ref.size - 1)
|
| 186 |
+
eq = ref[pos_c] == g
|
| 187 |
+
if eq.any():
|
| 188 |
+
# Fold per-task attribution into bit counts; no item text is ever reconstructed.
|
| 189 |
+
bits = refmask[pos_c[eq]]
|
| 190 |
+
for ti in range(len(REFS)):
|
| 191 |
+
c = int((bits & np.uint16(1 << ti)).astype(np.uint16).sum())
|
| 192 |
+
if c:
|
| 193 |
+
name = REFS[ti][0]
|
| 194 |
+
hits[name] = hits.get(name, 0) + c
|
| 195 |
+
for srcname in rec["sources"]:
|
| 196 |
+
d0 = hits_by_source.setdefault(srcname, {})
|
| 197 |
+
d0[name] = d0.get(name, 0) + c // max(1, len(rec["sources"]))
|
| 198 |
+
if docs % 100000 < 3000:
|
| 199 |
print(f" audited {docs:,} docs, {grams_seen:,} mix grams", flush=True)
|
| 200 |
|
| 201 |
total_overlap = sum(hits.values())
|
| 202 |
+
expected_fp = grams_seen * ref.size / float(1 << 64)
|
| 203 |
report = {
|
| 204 |
"k": K,
|
| 205 |
+
"fingerprint": "multiplicative rolling hash mod 2^64, P=0x100000001B3",
|
| 206 |
+
"reference_items_by_split": refcounts,
|
| 207 |
"reference_unreadable": errors,
|
| 208 |
+
"reference_grams_distinct": int(np.unique(ref).size),
|
| 209 |
+
"reference_grams_total": int(ref.size),
|
| 210 |
"shards_expected": len(man["shards"]),
|
| 211 |
"shards_missing": missing,
|
| 212 |
"mix_documents_audited": docs,
|
| 213 |
"mix_grams_audited": grams_seen,
|
| 214 |
+
"overlap_by_task": hits,
|
|
|
|
| 215 |
"overlap_by_source": hits_by_source,
|
| 216 |
+
"expected_false_positives": round(expected_fp, 3),
|
| 217 |
+
"overlap_total": total_overlap,
|
| 218 |
+
"overlap_exceeds_noise": total_overlap > max(10.0, 3.0 * expected_fp),
|
| 219 |
+
"note": ("counts only; no item text read or printed. Reference splits are train/validation/dev "
|
| 220 |
+
"per task -- never test, per §3.3. Overlap above the stated false-positive expectation "
|
| 221 |
+
"means the named source must be dropped or cleaned before Gate 2 is claimed."),
|
| 222 |
+
"COVERED_ALL_SHARDS": bool(man["shards"] and not missing),
|
| 223 |
+
"MEASURED": bool(docs > 0 and grams_seen > 0 and len(refcounts) >= 6 and not errors),
|
| 224 |
+
"AUDIT_PASSED": bool(man["shards"] and not missing and docs > 0 and grams_seen > 0
|
| 225 |
+
and len(refcounts) >= 6 and not errors
|
| 226 |
+
and total_overlap <= max(10.0, 3.0 * expected_fp)),
|
| 227 |
}
|
| 228 |
+
txt = json.dumps(report, indent=1, default=str)
|
| 229 |
print(f"VERDICT AUDIT_PASSED={report['AUDIT_PASSED']} docs={docs:,} overlap={total_overlap} "
|
| 230 |
+
f"expected_fp={expected_fp:.2f} missing_shards={len(missing)} refs_ok={len(refcounts)} "
|
| 231 |
+
f"refs_failed={len(errors)}")
|
| 232 |
print("AUDIT_JSON_BEGIN")
|
| 233 |
+
print(txt)
|
| 234 |
print("AUDIT_JSON_END")
|
| 235 |
+
if a.out:
|
| 236 |
+
with open(a.out, "w") as f:
|
| 237 |
+
f.write(txt)
|
| 238 |
raise SystemExit(0 if report["AUDIT_PASSED"] else (6 if report["MEASURED"] else 7))
|
| 239 |
|
| 240 |
|