| """Counter identifier memorization with randomized, semantically equivalent training schemas.""" |
| import argparse |
| import copy |
| import hashlib |
| import json |
| import random |
| import re |
| import shutil |
| import string |
| from pathlib import Path |
| from tinyquery.data import SQL_OPS,ddls,gold_sql,serialize,compact |
|
|
| CHUNKS=[] |
| PLAIN=False |
|
|
| def random_id(rng,prefix): |
| if PLAIN and prefix in ['t','g','c'] and rng.random()<.85: |
| from sqlglot import Tokenizer as SQLTokenizer |
| while True: |
| value=''.join(rng.choices(CHUNKS,k=rng.choice([3,4,5]))) if CHUNKS and rng.random()<.8 else ''.join(rng.choices(string.ascii_lowercase,k=rng.choice([6,7,8]))) |
| if len(value)>=6 and value.upper() not in SQLTokenizer.KEYWORDS: return value |
| if CHUNKS and rng.random()<.8: |
| return prefix+'_'+''.join(rng.choices(CHUNKS,k=rng.choice([2,3,4]))) |
| return prefix+'_'+''.join(rng.choices(string.ascii_lowercase,k=rng.choice([4,5,6,7]))) |
|
|
|
|
| def main(): |
| global PLAIN |
| p=argparse.ArgumentParser(); p.add_argument('--source',default='data/tinyquery');p.add_argument('--out',default='data/tinyquery-v2') |
| p.add_argument('--token-chunks',action='store_true'); p.add_argument('--seed',type=int,default=99032) |
| p.add_argument('--extra',help='Replay an earlier augmentation corpus, deduplicated by prompt') |
| p.add_argument('--plain-identifiers',action='store_true') |
| args=p.parse_args(); source=Path(args.source); out=Path(args.out); out.mkdir(parents=True,exist_ok=True) |
| PLAIN=args.plain_identifiers |
| for name in ['validation.jsonl','test.jsonl','manual.jsonl','tokenizer.json']: |
| shutil.copy2(source/name,out/name) |
| if args.token_chunks: |
| from tokenizers import Tokenizer |
| tokenizer=Tokenizer.from_file(str(source/'tokenizer.json')) |
| CHUNKS.extend(sorted({tokenizer.decode([i]) for i in range(tokenizer.get_vocab_size()) |
| if re.fullmatch('[a-z]{1,8}',tokenizer.decode([i]))})) |
| rng=random.Random(args.seed); counts={'original':0,'grounding':0,'discovery':0,'replay':0}; seen=set() |
| with (out/'train.jsonl').open('w') as stream: |
| def emit(r,kind): |
| key=hashlib.sha256(r['prompt'].encode()).hexdigest() |
| if key in seen: return |
| seen.add(key); counts[kind]+=1; stream.write(json.dumps(r,ensure_ascii=False)+'\n') |
| with (source/'train.jsonl').open() as original: |
| for line in original: emit(json.loads(line),'original') |
| with (source/'train-base.jsonl').open() as original: |
| for index,line in enumerate(original): |
| r=json.loads(line) |
| if re.search('[\u0600-\u06ff]',r['question']): continue |
| s=r['slots']; mapping={} |
| if r['operation'] in SQL_OPS and rng.random()<.85: |
| for field in ['table','parent']+(['column','numeric','category','date_column','foreign','label'] if rng.random()<.6 else []): |
| old=s[field]; s[field]=random_id(rng,{'table':'t','parent':'g'}.get(field,'c')) |
| mapping[old]=s[field] |
| if rng.random()<.5: |
| for field in ['value','other']: |
| old=str(s[field]); s[field]=mapping.get(old) or ''.join(rng.choices(string.ascii_letters,k=7)) |
| mapping[old]=s[field] |
| |
| pattern=re.compile(r'(?<![A-Za-z0-9_])('+ '|'.join(re.escape(k) for k in sorted(mapping,key=len,reverse=True))+r')(?![A-Za-z0-9_])') |
| r['question']=pattern.sub(lambda match:mapping[match.group()],r['question']) |
| r['context']['schema']=ddls(s) |
| key=next(k for k in ('sql','query') if k in r['target']['arguments']) |
| r['target']['arguments'][key]=gold_sql(r['operation'],s,r['backend']) |
| if r['target']['action']=='call': |
| old_target=r['target']['name'] |
| for t in r['context']['tools']: |
| old=t['name'] |
| |
| t['name']=random_id(rng,rng.choice(['tool','api','connector','fn'])) |
| if CHUNKS and rng.random()<.8: |
| suffix=old.rsplit('.',1)[-1].rsplit('__',1)[-1] |
| prefix=''.join(rng.choices(CHUNKS,k=rng.choice([2,3]))) |
| t['name']=prefix+rng.choice(['.','__','_'])+suffix |
| if old==old_target: r['target']['name']=t['name'] |
| if 'project_id' in r['target']['arguments']: |
| project=random_id(rng,'project') |
| r['context']['project_id']=project; r['target']['arguments']['project_id']=project |
| r['id']+=f'_grounded_{args.seed}'; r['response']=compact(r['target']); r['prompt']=serialize(r['context'],r['question']) |
| r['provenance']+='; randomized tool names, identifiers and literal values to train copying from input' |
| emit(r,'grounding') |
| with (source/'discovery-augmentation.jsonl').open() as extra: |
| for line in extra: emit(json.loads(line),'discovery') |
| if args.extra: |
| with Path(args.extra).open() as extra: |
| for line in extra: emit(json.loads(line),'replay') |
| (out/'grounding-stats.json').write_text(json.dumps(counts,indent=2)); print(counts) |
|
|
|
|
| if __name__=='__main__': main() |
|
|