"""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] # All source templates contain the explicit identifier/value slots being renamed. pattern=re.compile(r'(?