infosec-v1 / code /training /scripts /decontaminate.py
adhikjoshi's picture
Super-squash branch 'main' using huggingface_hub
994182c
Raw
History Blame Contribute Delete
12.1 kB
#!/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())