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