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())