media-profiling / scripts /diagnose_contamination.py
Claude
Improve FACTScore precision: negative assertion filtering, dense retrieval, LLM synthesis mode
61d96e4
Raw
History Blame Contribute Delete
7.11 kB
#!/usr/bin/env python3
"""Diagnostic: Quantify training data contamination in LLM mode.
Measures n-gram overlap between LLM-generated text and MBFC gold text
at the sentence level. High overlap indicates the LLM is reproducing
memorized MBFC content rather than performing independent analysis.
Compares overlap rates across modes (LLM vs System vs Hybrid) to
quantify how much each mode relies on memorized vs discovered facts.
Usage:
python scripts/diagnose_contamination.py [--results-dir results] [--model gpt-5-mini-2025-08-07]
"""
import argparse
import json
import logging
import os
import re
import sys
from collections import defaultdict
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s")
logger = logging.getLogger(__name__)
def extract_ngrams(text: str, n: int = 4) -> set[tuple[str, ...]]:
"""Extract word-level n-grams from text."""
words = re.findall(r'\b\w+\b', text.lower())
return {tuple(words[i:i+n]) for i in range(len(words) - n + 1)}
def sentence_overlap(gen_text: str, gold_text: str, ngram_size: int = 4) -> dict:
"""Compute sentence-level overlap metrics between generated and gold text."""
gen_sentences = [s.strip() for s in re.split(r'(?<=[.!?])\s+', gen_text) if s.strip()]
gold_sentences = [s.strip() for s in re.split(r'(?<=[.!?])\s+', gold_text) if s.strip()]
gold_ngrams = extract_ngrams(gold_text, ngram_size)
if not gold_ngrams:
return {"overlap_ratio": 0.0, "verbatim_sentences": 0, "total_sentences": len(gen_sentences)}
verbatim_count = 0
high_overlap_count = 0
for sent in gen_sentences:
sent_ngrams = extract_ngrams(sent, ngram_size)
if not sent_ngrams:
continue
overlap = len(sent_ngrams & gold_ngrams) / len(sent_ngrams)
if overlap > 0.8:
verbatim_count += 1
elif overlap > 0.5:
high_overlap_count += 1
gen_ngrams = extract_ngrams(gen_text, ngram_size)
overall_overlap = len(gen_ngrams & gold_ngrams) / len(gen_ngrams) if gen_ngrams else 0.0
return {
"overall_ngram_overlap": overall_overlap,
"verbatim_sentences": verbatim_count,
"high_overlap_sentences": high_overlap_count,
"total_sentences": len(gen_sentences),
"verbatim_ratio": verbatim_count / len(gen_sentences) if gen_sentences else 0.0,
}
def load_results(results_dir: str, mode: str, model: str) -> list[dict]:
path = os.path.join(results_dir, f"{model}_{mode}.jsonl")
if not os.path.exists(path):
return []
with open(path) as f:
return [json.loads(line) for line in f]
def load_dataset(data_path: str = "data/mbfc_benchmark.json") -> dict:
full_path = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), data_path)
with open(full_path) as f:
return {item["name"]: item for item in json.load(f)}
def build_text(raw: dict) -> str:
return " ".join(filter(None, [
raw.get('bias_category_description'),
raw.get('overall_summary'),
raw.get('analysis'),
raw.get('history'),
raw.get('ownership'),
])).strip()
def build_gold_text(item: dict) -> str:
return " ".join(filter(None, [
item.get('bias_category_description', ''),
item.get('overall_summary', ''),
item.get('analysis', ''),
item.get('history', ''),
item.get('ownership', ''),
])).strip()
def analyze_mode(results: list[dict], gold_data: dict, mode: str) -> dict:
"""Analyze contamination for a single mode."""
overlaps = []
for result in results:
name = result.get("name")
gold = gold_data.get(name)
if not gold:
continue
raw = result.get("raw_output", {})
gen_text = build_text(raw)
gold_text = build_gold_text(gold)
if not gen_text or not gold_text:
continue
metrics = sentence_overlap(gen_text, gold_text)
metrics["outlet"] = name
overlaps.append(metrics)
if not overlaps:
return {}
avg_overlap = sum(o["overall_ngram_overlap"] for o in overlaps) / len(overlaps)
avg_verbatim = sum(o["verbatim_ratio"] for o in overlaps) / len(overlaps)
total_verbatim = sum(o["verbatim_sentences"] for o in overlaps)
total_sentences = sum(o["total_sentences"] for o in overlaps)
return {
"mode": mode,
"outlets_analyzed": len(overlaps),
"avg_ngram_overlap": round(avg_overlap, 4),
"avg_verbatim_ratio": round(avg_verbatim, 4),
"total_verbatim_sentences": total_verbatim,
"total_sentences": total_sentences,
"per_outlet": sorted(overlaps, key=lambda x: x["overall_ngram_overlap"], reverse=True),
}
def main():
parser = argparse.ArgumentParser(description="Diagnose training data contamination")
parser.add_argument("--results-dir", default="results")
parser.add_argument("--model", default="gpt-5-mini-2025-08-07")
parser.add_argument("--modes", default="llm,system,hybrid,articles")
parser.add_argument("--dataset", default="data/mbfc_benchmark.json")
args = parser.parse_args()
gold_data = load_dataset(args.dataset)
modes = args.modes.split(",")
print(f"\n{'='*70}")
print(f"CONTAMINATION DIAGNOSTIC REPORT")
print(f"{'='*70}")
print(f"Model: {args.model}")
print(f"4-gram overlap between generated text and MBFC gold text")
print(f"Higher overlap = more likely memorized from training data")
print(f"{'='*70}\n")
all_results = {}
for mode in modes:
results = load_results(args.results_dir, mode, args.model)
if not results:
print(f" {mode}: No results found")
continue
analysis = analyze_mode(results, gold_data, mode)
if not analysis:
continue
all_results[mode] = analysis
print(f" {mode:12s}: avg_overlap={analysis['avg_ngram_overlap']:.3f} "
f"verbatim_ratio={analysis['avg_verbatim_ratio']:.3f} "
f"verbatim_sents={analysis['total_verbatim_sentences']}/{analysis['total_sentences']}")
print()
# Show per-outlet details for the mode with highest overlap
if all_results:
highest_mode = max(all_results, key=lambda m: all_results[m]["avg_ngram_overlap"])
print(f"Top 10 most contaminated outlets ({highest_mode} mode):")
for item in all_results[highest_mode]["per_outlet"][:10]:
print(f" {item['outlet']:30s}: overlap={item['overall_ngram_overlap']:.3f} "
f"verbatim={item['verbatim_sentences']}/{item['total_sentences']}")
# Save full results
output_path = os.path.join(args.results_dir, "diagnostic_contamination.json")
with open(output_path, "w") as f:
json.dump({mode: {k: v for k, v in data.items() if k != "per_outlet"}
for mode, data in all_results.items()}, f, indent=2)
print(f"\nSummary saved to {output_path}")
if __name__ == "__main__":
main()