File size: 3,588 Bytes
811d51e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
from __future__ import annotations
import argparse, json, random, re
from collections import Counter, defaultdict
from pathlib import Path

TASKS=['general_chat','writing','translation','summarization','research','coding','mathematics','document_analysis','high_stakes']
POLITE_PREFIXES=['Please ','Could you ','I need you to ','Help me to ']

def normalize(s): return re.sub(r'\s+',' ',s).strip()
def variants(text):
    base=normalize(text)
    out=[]
    if base:
        out.append(base)
        if base[0].islower(): out.append(base[0].upper()+base[1:])
        if not base.endswith(('?','.','!')): out.append(base+'.')
        for p in POLITE_PREFIXES:
            if not base.lower().startswith(('please ','could you ','i need you to ','help me to ')):
                out.append(p+base[0].lower()+base[1:])
    seen=[]
    for x in out:
        if x not in seen: seen.append(x)
    return seen

def main():
    ap=argparse.ArgumentParser()
    ap.add_argument('input',type=Path)
    ap.add_argument('--output',type=Path,default=Path('data/processed/training-balanced.jsonl'))
    ap.add_argument('--report',type=Path,default=Path('data/processed/training-balanced-report.json'))
    ap.add_argument('--seed',type=int,default=42)
    ap.add_argument('--confidence',type=float,default=.65)
    ap.add_argument('--target-per-task',type=int,default=800)
    ap.add_argument('--max-per-task',type=int,default=2000)
    args=ap.parse_args(); rng=random.Random(args.seed)
    rows=[json.loads(x) for x in args.input.read_text(encoding='utf-8').splitlines() if x.strip()]
    accepted=[]; rejected=[]
    for r in rows:
        fixed=r.get('label_method') in {'source_fixed','human_override'}
        if fixed or float(r.get('label_confidence',0))>=args.confidence:
            accepted.append(r)
        else: rejected.append(r)
    groups=defaultdict(list)
    for r in accepted: groups[r['task']].append(r)
    output=[]; augmented=Counter()
    for task in TASKS:
        group=groups[task]
        rng.shuffle(group)
        selected=group[:args.max_per_task]
        output.extend(selected)
        needed=max(0,args.target_per_task-len(selected))
        if needed and selected:
            pool=[]
            for r in selected:
                for i,v in enumerate(variants(r['text'])[1:],1):
                    n=dict(r); n['id']=f"{r['id']}:aug{i}"; n['text']=v
                    n['label_method']='deterministic_augmentation'; n['derived_from']=r['id']; n['label_confidence']=r.get('label_confidence',1.0)
                    pool.append(n)
            rng.shuffle(pool)
            take=pool[:needed]
            output.extend(take); augmented[task]+=len(take)
    rng.shuffle(output)
    args.output.parent.mkdir(parents=True,exist_ok=True)
    with args.output.open('w',encoding='utf-8') as f:
        for r in output: f.write(json.dumps(r,ensure_ascii=False)+'\n')
    report={
      'input_rows':len(rows),'accepted_rows':len(accepted),'rejected_low_confidence':len(rejected),
      'output_rows':len(output),'confidence_threshold':args.confidence,
      'by_task':dict(Counter(r['task'] for r in output)),
      'by_complexity':dict(Counter(r['complexity'] for r in output)),
      'by_label_method':dict(Counter(r.get('label_method','unknown') for r in output)),
      'augmented_by_task':dict(augmented),
      'warning':'Augmented rows are lexical variants, not independent human examples.'
    }
    args.report.write_text(json.dumps(report,indent=2),encoding='utf-8')
    print(json.dumps(report,indent=2))
if __name__=='__main__': main()