File size: 13,269 Bytes
803b5e8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
"""

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."""
    # Remove punctuation, lowercase, collapse whitespace
    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
    # Quick length check
    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', [])
    
    # Must have system, user, and at least 2 assistant turns
    if len(trace_msgs) < 4:
        return False, "too_few_messages"
    
    # Query must be substantial
    if len(query) < 15:
        return False, "query_too_short"
    
    # Check assistant turns have real content
    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"
        # Must contain at least one special token
        if not any(tok in content for tok in ['<|reasoning|>', '<|search|>', '<|evidence|>', '<|finish|>']):
            return False, "no_special_tokens"
    
    # Must have at least one search
    has_search = any('<|search|>' in m.get('content', '') for m in assistant_turns)
    if not has_search:
        return False, "no_search_action"
    
    # Must have evidence or finish
    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"
    
    # Check for reasoning density — at least one turn should have substantial reasoning
    max_reasoning_len = 0
    for turn in assistant_turns:
        content = turn.get('content', '')
        # Extract reasoning sections
        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),
    }
    
    # Phase 1: Exact dedup by query hash
    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
    
    # Phase 2: Fuzzy dedup by query similarity
    # Group by first word for efficiency
    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
        
        # Compare within group
        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)

    # 1. Load all traces
    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):,}")

    # 2. Quality filter
    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}")

    # 3. Contamination check
    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):,}")

    # 4. Deduplication
    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']:,}")

    # 5. Shuffle and write
    print(f"\n5. Writing to {args.output}...")
    rng = random.Random(args.seed)
    rng.shuffle(deduped_traces)
    
    # Remove internal fields
    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")
    
    # Summary
    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):,}")
    
    # Category distribution
    cat_counts = defaultdict(int)
    for trace in deduped_traces:
        # Try to infer category from query pattern
        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()