| """
|
| Deduplication, contamination check, and quality filter for SFT traces.
|
|
|
| Pipeline:
|
| 1. Load all traces from data/sft_traces_v2/ + existing sft_traces.jsonl
|
| 2. Deduplicate by query hash (exact match) and by fuzzy similarity (near-dupes)
|
| 3. Contamination check: remove traces whose queries appear in gold_traces.jsonl
|
| 4. Quality filter: remove traces that are too short, have empty reasoning,
|
| or have malformed structure
|
| 5. Shuffle and write final dataset
|
|
|
| Usage:
|
| python src/dedup_quality.py --input data/sft_traces_v2/ --output data/sft_traces_final.jsonl
|
| """
|
|
|
| import argparse
|
| import hashlib
|
| import json
|
| import os
|
| import random
|
| import re
|
| import sys
|
| from collections import defaultdict
|
| from difflib import SequenceMatcher
|
|
|
| PROJECT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
|
|
| def load_traces_from_dir(dir_path: str) -> list[dict]:
|
| """Load all .jsonl files from a directory."""
|
| traces = []
|
| if not os.path.exists(dir_path):
|
| return traces
|
| for fname in sorted(os.listdir(dir_path)):
|
| if fname.endswith('.jsonl'):
|
| fpath = os.path.join(dir_path, fname)
|
| with open(fpath, 'r', encoding='utf-8') as f:
|
| for line in f:
|
| try:
|
| trace = json.loads(line)
|
| if trace and 'query' in trace and 'trace' in trace:
|
| trace['_source'] = fname
|
| traces.append(trace)
|
| except json.JSONDecodeError:
|
| continue
|
| return traces
|
|
|
| def load_traces_from_file(fpath: str) -> list[dict]:
|
| """Load traces from a single .jsonl file."""
|
| traces = []
|
| if not os.path.exists(fpath):
|
| return traces
|
| with open(fpath, 'r', encoding='utf-8') as f:
|
| for line in f:
|
| try:
|
| trace = json.loads(line)
|
| if trace and 'query' in trace and 'trace' in trace:
|
| trace['_source'] = os.path.basename(fpath)
|
| traces.append(trace)
|
| except json.JSONDecodeError:
|
| continue
|
| return traces
|
|
|
| def query_hash(trace: dict) -> str:
|
| """Hash the query for exact dedup."""
|
| return hashlib.md5(trace['query'].strip().lower().encode()).hexdigest()
|
|
|
| def normalize_query(query: str) -> str:
|
| """Normalize query for fuzzy matching."""
|
|
|
| q = re.sub(r'[^\w\s]', '', query.lower())
|
| q = ' '.join(q.split())
|
| return q
|
|
|
| def fuzzy_similarity(q1: str, q2: str) -> float:
|
| """Compute similarity between two queries."""
|
| n1, n2 = normalize_query(q1), normalize_query(q2)
|
| if n1 == n2:
|
| return 1.0
|
|
|
| if abs(len(n1) - len(n2)) > max(len(n1), len(n2)) * 0.5:
|
| return 0.0
|
| return SequenceMatcher(None, n1, n2).ratio()
|
|
|
| def check_trace_quality(trace: dict) -> tuple[bool, str]:
|
| """Check if a trace meets quality standards.
|
|
|
| Returns (is_valid, reason_if_rejected)
|
| """
|
| query = trace.get('query', '')
|
| trace_msgs = trace.get('trace', [])
|
|
|
|
|
| if len(trace_msgs) < 4:
|
| return False, "too_few_messages"
|
|
|
|
|
| if len(query) < 15:
|
| return False, "query_too_short"
|
|
|
|
|
| assistant_turns = [m for m in trace_msgs if m['role'] == 'assistant']
|
| if len(assistant_turns) < 2:
|
| return False, "too_few_assistant_turns"
|
|
|
| for turn in assistant_turns:
|
| content = turn.get('content', '')
|
| if len(content) < 50:
|
| return False, "assistant_turn_too_short"
|
|
|
| if not any(tok in content for tok in ['<|reasoning|>', '<|search|>', '<|evidence|>', '<|finish|>']):
|
| return False, "no_special_tokens"
|
|
|
|
|
| has_search = any('<|search|>' in m.get('content', '') for m in assistant_turns)
|
| if not has_search:
|
| return False, "no_search_action"
|
|
|
|
|
| has_evidence = any('<|evidence|>' in m.get('content', '') for m in assistant_turns)
|
| has_finish = any('<|finish|>' in m.get('content', '') for m in assistant_turns)
|
| if not has_evidence and not has_finish:
|
| return False, "no_evidence_or_finish"
|
|
|
|
|
| max_reasoning_len = 0
|
| for turn in assistant_turns:
|
| content = turn.get('content', '')
|
|
|
| reasoning_sections = re.findall(r'<\|reasoning\|>(.*?)<\|end\|>', content, re.DOTALL)
|
| for r in reasoning_sections:
|
| max_reasoning_len = max(max_reasoning_len, len(r.strip()))
|
|
|
| if max_reasoning_len < 30:
|
| return False, "reasoning_too_thin"
|
|
|
| return True, "ok"
|
|
|
| def deduplicate(traces: list[dict], similarity_threshold: float = 0.85) -> tuple[list[dict], dict]:
|
| """Remove duplicate and near-duplicate traces.
|
|
|
| Returns (deduplicated_traces, stats)
|
| """
|
| stats = {
|
| 'exact_dups_removed': 0,
|
| 'fuzzy_dups_removed': 0,
|
| 'total_input': len(traces),
|
| }
|
|
|
|
|
| seen_hashes = set()
|
| exact_deduped = []
|
| for trace in traces:
|
| h = query_hash(trace)
|
| if h not in seen_hashes:
|
| seen_hashes.add(h)
|
| exact_deduped.append(trace)
|
| else:
|
| stats['exact_dups_removed'] += 1
|
|
|
|
|
|
|
| groups = defaultdict(list)
|
| for trace in exact_deduped:
|
| first_word = normalize_query(trace['query']).split()[0] if normalize_query(trace['query']).split() else ''
|
| groups[first_word].append(trace)
|
|
|
| fuzzy_deduped = []
|
| for first_word, group in groups.items():
|
| if len(group) == 1:
|
| fuzzy_deduped.extend(group)
|
| continue
|
|
|
|
|
| kept = []
|
| for trace in group:
|
| is_dup = False
|
| for kept_trace in kept:
|
| sim = fuzzy_similarity(trace['query'], kept_trace['query'])
|
| if sim >= similarity_threshold:
|
| is_dup = True
|
| stats['fuzzy_dups_removed'] += 1
|
| break
|
| if not is_dup:
|
| kept.append(trace)
|
| fuzzy_deduped.extend(kept)
|
|
|
| stats['total_output'] = len(fuzzy_deduped)
|
| return fuzzy_deduped, stats
|
|
|
| def check_contamination(traces: list[dict], gold_traces: list[dict]) -> tuple[list[dict], dict]:
|
| """Remove traces whose queries match gold trace queries.
|
|
|
| Returns (clean_traces, stats)
|
| """
|
| gold_queries = set()
|
| for gt in gold_traces:
|
| gold_queries.add(normalize_query(gt['query']))
|
|
|
| clean = []
|
| removed = 0
|
| for trace in traces:
|
| nq = normalize_query(trace['query'])
|
| if nq in gold_queries:
|
| removed += 1
|
| else:
|
| clean.append(trace)
|
|
|
| return clean, {'contamination_removed': removed, 'gold_queries': len(gold_queries)}
|
|
|
| def main():
|
| parser = argparse.ArgumentParser(description="Dedup and quality filter SFT traces")
|
| parser.add_argument("--input", type=str, default=os.path.join(PROJECT_DIR, "data", "sft_traces_v2"),
|
| help="Input directory with .jsonl files")
|
| parser.add_argument("--existing", type=str, default=os.path.join(PROJECT_DIR, "data", "sft_traces.jsonl"),
|
| help="Existing traces to merge with")
|
| parser.add_argument("--gold", type=str, default=os.path.join(PROJECT_DIR, "data", "gold_traces.jsonl"),
|
| help="Gold traces for contamination check")
|
| parser.add_argument("--output", type=str, default=os.path.join(PROJECT_DIR, "data", "sft_traces_final.jsonl"),
|
| help="Output file")
|
| parser.add_argument("--seed", type=int, default=42)
|
| parser.add_argument("--similarity", type=float, default=0.85,
|
| help="Fuzzy dedup similarity threshold")
|
| args = parser.parse_args()
|
|
|
| print("=" * 60)
|
| print("SFT TRACE DEDUPLICATION & QUALITY PIPELINE")
|
| print("=" * 60)
|
|
|
|
|
| print("\n1. Loading traces...")
|
| new_traces = load_traces_from_dir(args.input)
|
| existing_traces = load_traces_from_file(args.existing)
|
| gold_traces = load_traces_from_file(args.gold)
|
|
|
| print(f" New traces: {len(new_traces):,}")
|
| print(f" Existing traces: {len(existing_traces):,}")
|
| print(f" Gold traces: {len(gold_traces):,}")
|
|
|
| all_traces = new_traces + existing_traces
|
| print(f" Total to process: {len(all_traces):,}")
|
|
|
|
|
| print("\n2. Quality filtering...")
|
| quality_stats = defaultdict(int)
|
| quality_passed = []
|
| for trace in all_traces:
|
| is_valid, reason = check_trace_quality(trace)
|
| if is_valid:
|
| quality_passed.append(trace)
|
| else:
|
| quality_stats[reason] += 1
|
|
|
| print(f" Passed: {len(quality_passed):,}")
|
| print(f" Rejected: {sum(quality_stats.values()):,}")
|
| for reason, count in sorted(quality_stats.items(), key=lambda x: -x[1]):
|
| print(f" {reason}: {count}")
|
|
|
|
|
| print("\n3. Contamination check (vs gold traces)...")
|
| clean_traces, contam_stats = check_contamination(quality_passed, gold_traces)
|
| print(f" Removed: {contam_stats['contamination_removed']}")
|
| print(f" Remaining: {len(clean_traces):,}")
|
|
|
|
|
| print("\n4. Deduplication...")
|
| deduped_traces, dedup_stats = deduplicate(clean_traces, args.similarity)
|
| print(f" Exact dups removed: {dedup_stats['exact_dups_removed']:,}")
|
| print(f" Fuzzy dups removed: {dedup_stats['fuzzy_dups_removed']:,}")
|
| print(f" Final count: {dedup_stats['total_output']:,}")
|
|
|
|
|
| print(f"\n5. Writing to {args.output}...")
|
| rng = random.Random(args.seed)
|
| rng.shuffle(deduped_traces)
|
|
|
|
|
| for trace in deduped_traces:
|
| trace.pop('_source', None)
|
|
|
| with open(args.output, 'w', encoding='utf-8') as f:
|
| for trace in deduped_traces:
|
| f.write(json.dumps(trace, ensure_ascii=False) + '\n')
|
|
|
| print(f" Written: {len(deduped_traces):,} traces")
|
|
|
|
|
| print("\n" + "=" * 60)
|
| print("SUMMARY")
|
| print("=" * 60)
|
| print(f" Input traces: {len(all_traces):,}")
|
| print(f" Quality rejected: {sum(quality_stats.values()):,}")
|
| print(f" Contamination removed: {contam_stats['contamination_removed']}")
|
| print(f" Exact dups removed: {dedup_stats['exact_dups_removed']:,}")
|
| print(f" Fuzzy dups removed: {dedup_stats['fuzzy_dups_removed']:,}")
|
| print(f" Final dataset: {len(deduped_traces):,}")
|
|
|
|
|
| cat_counts = defaultdict(int)
|
| for trace in deduped_traces:
|
|
|
| q = trace['query'].lower()
|
| if any(w in q for w in ['walk me through', 'implementation', 'step by step', 'control flow']):
|
| cat_counts['implementation'] += 1
|
| elif any(w in q for w in ['trace how data', 'data flow', 'data path', 'interact']):
|
| cat_counts['cross_file'] += 1
|
| elif any(w in q for w in ['architecture', 'module', 'structure of', 'map out']):
|
| cat_counts['architecture'] += 1
|
| elif any(w in q for w in ['where is', 'used across', 'usage', 'called from']):
|
| cat_counts['usage'] += 1
|
| elif any(w in q for w in ['error', 'failure', 'fail', 'debug']):
|
| cat_counts['error/debug'] += 1
|
| elif any(w in q for w in ['api', 'contract', 'interface', 'parameters']):
|
| cat_counts['api'] += 1
|
| elif any(w in q for w in ['depend', 'dependency', 'blast radius', 'impact']):
|
| cat_counts['dependency/impact'] += 1
|
| elif any(w in q for w in ['fields', 'data structure', 'layout', 'memory']):
|
| cat_counts['data_structure'] += 1
|
| elif any(w in q for w in ['compare', 'tradeoff', 'vs', 'contrast']):
|
| cat_counts['comparison'] += 1
|
| elif any(w in q for w in ['security', 'validation', 'vulnerability']):
|
| cat_counts['security'] += 1
|
| elif any(w in q for w in ['performance', 'bottleneck', 'hot', 'optimize']):
|
| cat_counts['performance'] += 1
|
| elif any(w in q for w in ['pattern', 'design']):
|
| cat_counts['design_patterns'] += 1
|
| else:
|
| cat_counts['other'] += 1
|
|
|
| print(f"\n Category distribution:")
|
| for cat, count in sorted(cat_counts.items(), key=lambda x: -x[1]):
|
| print(f" {cat:20s}: {count:5d}")
|
|
|
|
|
| if __name__ == "__main__":
|
| main()
|
|
|