File size: 12,076 Bytes
994182c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 | #!/usr/bin/env python3
"""Decontaminate the training mix against frozen eval sets.
Removes any training row that overlaps an evaluation item, so CyberGym /
vuln-detection / knowledge-MCQ numbers measure generalization, not memorization.
Methods (in order of cost):
1. 13-gram collision -- word-level n-gram inverted index over the eval text;
a train row sharing >= --min-shared-ngrams 13-grams
with any eval item is contaminated (GPT-3/Llama style).
2. fuzzy whole-text -- difflib ratio for short rows that have no 13-grams,
flagged when ratio >= --fuzzy-threshold.
3. embedding (opt-in) -- cosine match via sentence-transformers if installed
and --use-embedding is set; otherwise skipped with a note.
Eval sources can be:
--eval-jsonl PATH JSONL whose user/assistant text (or a `text` field) is the eval item
--eval-text PATH plain text / id-per-line file (e.g. CyberGym frozen task ids,
or raw code snippets) -- each non-empty line is one eval item
Outputs:
--output-clean PATH train JSONL with contaminated rows removed (required for training)
--report PATH markdown decontamination report (the gate artifact)
Stdlib only, so it runs anywhere.
Example:
python training/scripts/decontaminate.py \
--train data/processed/stage1.ready.normalized.jsonl \
--eval-jsonl data/eval/vuln_detection_test.jsonl \
--eval-jsonl data/eval/knowledge_mcq.jsonl \
--eval-text reports/cybergym/frozen_level1_baseline_tasks.txt \
--output-clean data/processed/stage1.decontam.jsonl \
--report data/decontam/stage1_decontam_report.md
"""
from __future__ import annotations
import argparse
import json
import re
from difflib import SequenceMatcher
from pathlib import Path
from typing import Any, Iterable
WORD_RE = re.compile(r"[A-Za-z0-9_]+")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
parser.add_argument("--train", required=True, help="Training JSONL (rows with `messages` or `text`).")
parser.add_argument("--eval-jsonl", action="append", default=[], help="Eval JSONL source (repeatable).")
parser.add_argument("--eval-text", action="append", default=[], help="Eval text/id-per-line source (repeatable).")
parser.add_argument("--output-clean", required=True)
parser.add_argument("--report", required=True)
parser.add_argument("--n", type=int, default=13, help="N-gram size (default 13).")
parser.add_argument("--min-shared-ngrams", type=int, default=1,
help="Min shared n-grams to flag a row (default 1 = any collision).")
parser.add_argument("--fuzzy-threshold", type=float, default=0.90,
help="difflib ratio for short rows without n-grams (default 0.90).")
parser.add_argument("--fuzzy-max-eval", type=int, default=2000,
help="Cap eval items scanned per short row (perf guard).")
parser.add_argument("--use-embedding", action="store_true",
help="Additionally use sentence-transformers cosine match if available.")
parser.add_argument("--embedding-threshold", type=float, default=0.95)
return parser.parse_args()
def read_jsonl(path: Path) -> Iterable[dict[str, Any]]:
with path.open("r", encoding="utf-8") as fh:
for line_no, line in enumerate(fh, start=1):
line = line.strip()
if not line:
continue
try:
row = json.loads(line)
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid JSON in {path}:{line_no}: {exc}") from exc
if isinstance(row, dict):
yield row
def row_text(row: dict[str, Any]) -> str:
if isinstance(row.get("messages"), list):
return " ".join(
str(m.get("content", "")) for m in row["messages"] if m.get("role") in {"user", "assistant"}
)
if isinstance(row.get("text"), str):
return row["text"]
# eval JSONL may use other field names
parts = [str(row.get(k, "")) for k in ("question", "prompt", "input", "func", "code", "user", "assistant", "output")]
return " ".join(p for p in parts if p)
def tokens(text: str) -> list[str]:
return WORD_RE.findall(text.lower())
def ngrams(toks: list[str], n: int) -> set[str]:
if len(toks) < n:
return set()
return {" ".join(toks[i : i + n]) for i in range(len(toks) - n + 1)}
def load_eval_items(eval_jsonl: list[str], eval_text: list[str]) -> list[dict[str, Any]]:
items: list[dict[str, Any]] = []
for path in eval_jsonl:
p = Path(path)
if not p.is_file():
items.append({"_missing": str(p)})
continue
for i, row in enumerate(read_jsonl(p)):
items.append({"source": path, "id": row.get("id", f"{path}:{i}"), "text": row_text(row)})
for path in eval_text:
p = Path(path)
if not p.is_file():
items.append({"_missing": str(p)})
continue
for i, line in enumerate(p.read_text(encoding="utf-8").splitlines()):
line = line.strip()
if line and not line.startswith("#"):
items.append({"source": path, "id": f"{path}:{i}", "text": line})
return items
def build_index(eval_items: list[dict[str, Any]], n: int) -> tuple[dict[str, set[int]], list[set[str]], list[str]]:
"""Return (ngram -> eval-idx set, per-eval ngram sets, per-eval short text)."""
index: dict[str, set[int]] = {}
eval_ngrams: list[set[str]] = []
eval_short: list[str] = []
for idx, item in enumerate(eval_items):
text = item.get("text", "")
toks = tokens(text)
grams = ngrams(toks, n)
eval_ngrams.append(grams)
eval_short.append(text if len(toks) < n else "")
for g in grams:
index.setdefault(g, set()).add(idx)
return index, eval_ngrams, eval_short
def try_load_embedder(name: str = "all-MiniLM-L6-v2"):
try:
from sentence_transformers import SentenceTransformer # type: ignore
return SentenceTransformer(name)
except Exception:
return None
def main() -> int:
args = parse_args()
train_path = Path(args.train)
clean_path = Path(args.output_clean)
report_path = Path(args.report)
clean_path.parent.mkdir(parents=True, exist_ok=True)
report_path.parent.mkdir(parents=True, exist_ok=True)
eval_items = load_eval_items(args.eval_jsonl, args.eval_text)
missing = [i["_missing"] for i in eval_items if "_missing" in i]
eval_items = [i for i in eval_items if "_missing" not in i]
index, eval_ngrams, eval_short = build_index(eval_items, args.n)
short_eval_idx = [i for i, s in enumerate(eval_short) if s]
embedder = None
eval_embeddings = None
embed_note = "disabled"
if args.use_embedding:
embedder = try_load_embedder()
if embedder is None:
embed_note = "requested but sentence-transformers unavailable -> skipped"
else:
eval_embeddings = embedder.encode([i["text"] for i in eval_items], normalize_embeddings=True)
embed_note = "enabled (all-MiniLM-L6-v2)"
total = 0
kept = 0
flagged: list[dict[str, Any]] = []
reasons = {"ngram": 0, "fuzzy": 0, "embedding": 0}
with clean_path.open("w", encoding="utf-8") as out:
for row in read_jsonl(train_path):
total += 1
text = row_text(row)
toks = tokens(text)
grams = ngrams(toks, args.n)
hit_eval = None
reason = None
if grams:
counts: dict[int, int] = {}
for g in grams:
for eidx in index.get(g, ()): # eval items sharing this n-gram
counts[eidx] = counts.get(eidx, 0) + 1
if counts:
best = max(counts, key=counts.get)
if counts[best] >= args.min_shared_ngrams:
hit_eval, reason = best, "ngram"
reasons["ngram"] += 1
else:
# short row: fuzzy compare against short eval items
for eidx in short_eval_idx[: args.fuzzy_max_eval]:
ratio = SequenceMatcher(None, text, eval_short[eidx]).ratio()
if ratio >= args.fuzzy_threshold:
hit_eval, reason = eidx, "fuzzy"
reasons["fuzzy"] += 1
break
if hit_eval is None and embedder is not None and eval_embeddings is not None and text.strip():
import numpy as np # type: ignore
vec = embedder.encode([text], normalize_embeddings=True)[0]
sims = np.asarray(eval_embeddings) @ np.asarray(vec)
top = int(sims.argmax())
if float(sims[top]) >= args.embedding_threshold:
hit_eval, reason = top, "embedding"
reasons["embedding"] += 1
if hit_eval is not None:
flagged.append(
{
"train_id": row.get("id", "?"),
"reason": reason,
"eval_id": eval_items[hit_eval].get("id"),
"eval_source": eval_items[hit_eval].get("source"),
}
)
continue
out.write(json.dumps(row, ensure_ascii=False, sort_keys=True) + "\n")
kept += 1
_write_report(report_path, args, total, kept, flagged, reasons, eval_items, missing, embed_note)
print(json.dumps(
{
"train_rows": total,
"kept": kept,
"removed": len(flagged),
"reasons": reasons,
"eval_items": len(eval_items),
"missing_eval_sources": missing,
"clean": str(clean_path),
"report": str(report_path),
},
indent=2,
))
# Non-zero exit if eval sources were declared but missing (gate should not pass silently).
return 3 if missing else 0
def _write_report(path, args, total, kept, flagged, reasons, eval_items, missing, embed_note) -> None:
lines = [
"# Decontamination Report",
"",
f"- Train file: `{args.train}`",
f"- Eval items: {len(eval_items)} (from {len(args.eval_jsonl)} jsonl + {len(args.eval_text)} text sources)",
f"- N-gram size: {args.n}; min shared to flag: {args.min_shared_ngrams}",
f"- Fuzzy threshold (short rows): {args.fuzzy_threshold}",
f"- Embedding match: {embed_note}",
"",
f"- Train rows in: **{total}**",
f"- Kept (clean): **{kept}**",
f"- Removed (contaminated): **{len(flagged)}**",
f" - by n-gram: {reasons['ngram']}",
f" - by fuzzy: {reasons['fuzzy']}",
f" - by embedding: {reasons['embedding']}",
f"- Clean output: `{args.output_clean}`",
"",
]
if missing:
lines.append("## ⚠️ Missing eval sources (gate must not pass until resolved)")
for m in missing:
lines.append(f"- {m}")
lines.append("")
if flagged:
lines.append("## Sample of removed rows (first 50)")
lines.append("")
lines.append("| train_id | reason | eval_id | eval_source |")
lines.append("|---|---|---|---|")
for f in flagged[:50]:
lines.append(f"| {f['train_id']} | {f['reason']} | {f['eval_id']} | {f['eval_source']} |")
if len(flagged) > 50:
lines.append("")
lines.append(f"... and {len(flagged) - 50} more.")
else:
lines.append("No contamination detected against the provided eval sources.")
lines.append("")
Path(path).write_text("\n".join(lines) + "\n", encoding="utf-8")
if __name__ == "__main__":
raise SystemExit(main())
|