TinyQuery-140M / tinyquery /grounding_data.py
karmx's picture
Release TinyQuery 139.7M from scratch with frozen weights, reproducible Mac evaluations and runtime source
b296ad4 verified
Raw
History Blame Contribute Delete
5.63 kB
"""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'(?<![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']
# A tool's description/schema determines its purpose; names alone carry no reliable label.
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()