Cion-lab commited on
Commit
66e6f6e
·
verified ·
1 Parent(s): a477d8d

audit: vectorised rolling-hash grams + sorted-array lookup; the Python-set version could not finish 1.2B tokens

Browse files
Files changed (1) hide show
  1. 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 used as the reference, and why it is broader than the minimum:
10
- # for most of these tasks the harness draws its few-shot exemplars from the TRAIN split (ARC, HellaSwag,
11
- # PIQA, Winogrande) or from MMLU's `dev`, so overlap with those is contamination by construction -- a
12
- # mix that memorises an exemplar is a mix that inflates the score. So the reference set is
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
- # Method: for each reference document, tokenise with the frozen tokenizer and hash every 13-token window
17
- # into a set of uint64. For each mix document, do the same and count hits. 13-grams match the documented
18
- # practice of the decontaminated math source in the mix, so the threshold is comparable to published
19
- # hygiene rather than something invented here.
 
 
 
 
 
 
 
 
 
20
  #
21
- # Memory: the reference side is the small one -- a few hundred thousand short items x ~300 grams x 8 B,
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 array
34
 
35
  K = 13
36
- # task -> (repo, config(s), splits considered reference material; 'test' deliberately absent)
 
 
 
 
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 grams64(ids):
51
- """Set of uint64 hashes over K-length token windows. Hashing ids rather than characters means the
52
- comparison is in the same space as the mix, and no human-readable string is ever formed."""
53
- if len(ids) < K:
54
- return {struct.unpack("<Q", hashlib.sha256(bytes(array.array("H", ids))).digest()[:8])[0]} \
55
- if ids else set()
56
- out = set()
57
- for i in range(len(ids) - K + 1):
58
- w = ids[i:i + K]
59
- out.add(struct.unpack("<Q", hashlib.sha256(
60
- memoryview(bytes(array.array("H", w)))).digest()[:8])[0])
61
- return out
62
-
63
-
64
- def flatten(v, acc):
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
65
  if isinstance(v, str):
66
- acc.append(v)
67
  elif isinstance(v, dict):
68
  for x in v.values():
69
- flatten(x, acc)
70
  elif isinstance(v, (list, tuple)):
71
  for x in v:
72
- flatten(x, acc)
73
 
74
 
75
  def build_reference(tk, max_items_per_split):
76
- """Return {ref_key: gram_set}. Counts for the report are derived from these sets afterwards, so the
77
- reference corpus is streamed exactly once."""
78
  import datasets
79
 
80
- sets = {}
81
- for task, repo, configs, splits in REFS:
 
 
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
- sets[key] = {"error": f"{type(e).__name__}: {str(e)[:140]}"}
90
  print(f" ref {key}: UNREADABLE {type(e).__name__}", flush=True)
91
  continue
92
- g, n = set(), 0
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
- g |= grams64(tk.encode(txt, add_special_tokens=False).ids)
105
- sets[key] = g
106
- print(f" ref {key}: {n} items -> {len(g):,} grams", flush=True)
107
- return sets
 
 
 
 
 
 
 
 
 
 
 
 
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 sets (only counts ever leave this process)", flush=True)
122
- ref_sets = build_reference(tk, a.max_items_per_split)
123
- errors = {k: v["error"] for k, v in ref_sets.items() if isinstance(v, dict)}
124
- ref_sets = {k: v for k, v in ref_sets.items() if isinstance(v, set)}
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
- hits = {k: 0 for k in ref_sets}
 
 
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 that is not on disk is not "clean", it is unmeasured. Silence here would let an
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
- g = grams64(list(d))
144
- if not g:
 
145
  continue
146
- grams_seen += len(g)
147
- for k, s in ref_sets.items():
148
- ov = len(g & s)
149
- if ov:
150
- hits[k] += ov
151
- for srcname in rec["sources"]:
152
- hits_by_source.setdefault(srcname, {}).setdefault(k, 0)
153
- hits_by_source[srcname][k] += ov // max(1, len(rec["sources"]))
154
- if docs % 200000 < 5000:
 
 
 
 
 
 
 
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
- "reference": {k: {"distinct_grams": len(v)} for k, v in ref_sets.items()},
 
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
- "overlap_grams_total_by_ref": hits,
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
- # Composite, so Gate 2 cannot be claimed on an audit that measured nothing. Every clause here
170
- # has been observed to fail independently: unreadable refs, absent shards, zero docs.
171
- "COVERED_ALL_SHARDS": not missing and len(man["shards"]) > 0,
172
- "MEASURED": docs > 0 and grams_seen > 0 and len(ref_sets) >= 6 and not errors,
173
- "CLEAN": total_overlap == 0,
174
- "AUDIT_PASSED": bool(not missing and man["shards"] and docs > 0 and grams_seen > 0
175
- and len(ref_sets) >= 6 and not errors and total_overlap == 0),
176
- "note": ("counts only; no item text was read or printed. Reference splits are train/validation/dev "
177
- "per task -- never test, per §3.3. Any source with a non-zero overlap here must be dropped "
178
- "or cleaned before Gate 2 is claimed."),
 
179
  }
 
180
  print(f"VERDICT AUDIT_PASSED={report['AUDIT_PASSED']} docs={docs:,} overlap={total_overlap} "
181
- f"missing_shards={len(missing)} refs_ok={len(ref_sets)} refs_failed={len(errors)}")
 
182
  print("AUDIT_JSON_BEGIN")
183
- print(json.dumps(report, indent=1, default=str))
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