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
16 kB
"""Build fictitious tool/SQL data with semantic labels and grouped held-out domains."""
import argparse
import hashlib
import json
import random
import sqlite3
from pathlib import Path
from tinyquery.recipes import RECIPES, LANGUAGES
DOMAINS = {
'train': ['orders','products','employees','books','payments','students','invoices','tickets',
'shipments','courses','devices','expenses','appointments','subscriptions','accounts','visits'],
'validation': ['rentals','repairs','enrollments','deliveries'],
'test': ['specimens','exhibits','reservations','inspections'],
}
TEXT_COLS=['name','title','description','customer_name','item_name','display_name']
NUM_COLS=['amount','price','salary','score','cost','quantity','balance','rating']
CAT_COLS=['category','city','status','department','region','kind']
DATE_COLS=['created_at','order_date','start_date','recorded_at']
VALUES=['Delhi','Mumbai','Pune','active','pending','Books','Electronics','North','South',"O'Reilly"]
SQL_OPS=list(RECIPES)[:33]
def compact(value): return json.dumps(value,ensure_ascii=False,separators=(',',':'))
def serialize(context,question):
return '<|context|>'+compact(context)+'\n<|user|>'+question+'\n<|assistant|>'
def obj_schema(properties,required=None):
return {'type':'object','properties':properties,'required':list(properties) if required is None else required,
'additionalProperties':False}
def tool(name,description,properties,required=None):
return {'name':name,'description':description,'inputSchema':obj_schema(properties,required)}
def make_tools(backend,rng,split,project_id):
prefix=rng.choice(['','db.','database.']) if split=='train' else ('warehouse.' if split=='test' else 'store.')
string={'type':'string'}
query_key='query' if backend=='supabase' else rng.choice(['sql','query'])
query_name=prefix+('execute_sql' if backend=='supabase' else rng.choice(['query','run_query','execute_sql']))
scoped=backend!='supabase' or rng.random()<0.5
project={} if scoped else {'project_id':string}
tools=[tool(query_name,'Execute a read-only '+('PostgreSQL' if backend=='supabase' else 'MySQL')+' SELECT query.',
{**project,query_key:string}),
tool(prefix+'list_tables','List database tables and their columns.',
{**project,'schemas':{'type':'array','items':string}},list(project))]
if backend=='mysql':
tools.append(tool(prefix+'describe_table','Inspect columns and types for one table.',
{**project,'table':string}))
return tools,query_key,project,scoped
def schema_for(domain,index):
h=int(hashlib.sha256(domain.encode()).hexdigest()[:8],16)
return {'table':domain,'column':TEXT_COLS[(h+index)%len(TEXT_COLS)],
'numeric':NUM_COLS[(h+index)%len(NUM_COLS)],'category':CAT_COLS[(h+index)%len(CAT_COLS)],
'date_column':DATE_COLS[(h+index)%len(DATE_COLS)],'parent':domain+'_groups',
'foreign':'group_id','label':'label'}
def ddls(s):
return [f"CREATE TABLE {s['parent']} (id INTEGER PRIMARY KEY, {s['label']} TEXT);",
f"CREATE TABLE {s['table']} (id INTEGER PRIMARY KEY, {s['column']} TEXT, {s['numeric']} REAL, "
f"{s['category']} TEXT, {s['date_column']} DATE, {s['foreign']} INTEGER REFERENCES {s['parent']}(id));"]
def quote(value): return "'"+str(value).replace("'","''")+"'"
def gold_sql(op,s,backend):
t,c,n,g,d=s['table'],s['column'],s['numeric'],s['category'],s['date_column']
v,o=quote(s['value']),quote(s['other'])
x,y,k=s['number'],s['upper'],s['limit']
base=f'SELECT * FROM {t}'
queries={
'all':base,'project':f'SELECT {c} FROM {t}',
'eq':f'{base} WHERE {g} = {v}','project_eq':f'SELECT {c} FROM {t} WHERE {g} = {v}',
'gt':f'{base} WHERE {n} > {x}','lt':f'{base} WHERE {n} < {x}',
'gte':f'{base} WHERE {n} >= {x}','lte':f'{base} WHERE {n} <= {x}',
'between':f'{base} WHERE {n} BETWEEN {x} AND {y}',
'and':f'{base} WHERE {g} = {v} AND {n} > {x}',
'or':f'{base} WHERE {g} IN ({v}, {o})',
'null':f'{base} WHERE {c} IS NULL','not_null':f'{base} WHERE {c} IS NOT NULL',
'count':f'SELECT COUNT(*) FROM {t}','count_eq':f'SELECT COUNT(*) FROM {t} WHERE {g} = {v}',
'sum':f'SELECT SUM({n}) FROM {t}','avg':f'SELECT AVG({n}) FROM {t}',
'max':f'SELECT MAX({n}) FROM {t}','min':f'SELECT MIN({n}) FROM {t}',
'distinct':f'SELECT DISTINCT {g} FROM {t}',
'sort_desc':f'{base} ORDER BY {n} DESC','sort_asc':f'{base} ORDER BY {n} ASC',
'top':f'{base} ORDER BY {n} DESC LIMIT {k}','bottom':f'{base} ORDER BY {n} ASC LIMIT {k}',
'group_count':f'SELECT {g}, COUNT(*) FROM {t} GROUP BY {g}',
'group_sum':f'SELECT {g}, SUM({n}) FROM {t} GROUP BY {g}',
'having':f'SELECT {g}, COUNT(*) FROM {t} GROUP BY {g} HAVING COUNT(*) > {x}',
'date_after':f'{base} WHERE {d} >= '+quote(s['date']),
'year':f'{base} WHERE '+(f'YEAR({d})' if backend=='mysql' else f'EXTRACT(YEAR FROM {d})')+f" = {s['year']}",
'month':f'{base} WHERE '+(f'MONTH({d})' if backend=='mysql' else f'EXTRACT(MONTH FROM {d})')+f" = {s['month']}",
'contains':f'{base} WHERE LOWER({c}) LIKE LOWER('+quote('%'+s['value']+'%')+')',
'join':f"SELECT {t}.{c}, {s['parent']}.{s['label']} FROM {t} JOIN {s['parent']} ON {t}.{s['foreign']} = {s['parent']}.id",
'join_filter':f"SELECT {t}.{c} FROM {t} JOIN {s['parent']} ON {t}.{s['foreign']} = {s['parent']}.id WHERE {s['parent']}.{s['label']} = {v}",
}
return queries[op]+';'
def fixture(s,seed=0):
rng=random.Random(seed)
db=sqlite3.connect(':memory:')
db.create_function('YEAR',1,lambda d: None if d is None else int(str(d)[:4]))
db.create_function('MONTH',1,lambda d: None if d is None else int(str(d)[5:7]))
for ddl in ddls(s): db.execute(ddl)
db.executemany(f"INSERT INTO {s['parent']} VALUES (?,?)",[(1,s['value']),(2,s['other']),(3,'West')])
vals=[s['number']-1,s['number'],s['number']+1,s['upper'],s['upper']+1]
rows=[]
for i in range(24):
rows.append((i+1, None if i%9==0 else rng.choice(['Asha','Ravi','item '+s['value'],'Other']),
(vals[i%len(vals)] if i<6 else rng.choice(vals)+i*0.001) if i%6 else None,
rng.choice([s['value'],s['other'],'Other',None]),
rng.choice([s['date'],f"{s['year']}-01-01",'2022-07-18',None]),rng.choice([1,2,3,None])))
db.executemany(f"INSERT INTO {s['table']} VALUES (?,?,?,?,?,?)",rows)
db.execute('PRAGMA query_only=ON')
return db
def validate_sql(sql,backend,s):
import sqlglot
from sqlglot import exp
parsed=sqlglot.parse(sql,read='postgres' if backend=='supabase' else 'mysql')
if len(parsed)!=1 or not isinstance(parsed[0],exp.Select): raise ValueError('Not one SELECT')
result=[]
for seed in (1,7):
db=fixture(s,seed)
translated=sqlite_sql(parsed[0])
try: result.append(db.execute(translated).fetchall())
except Exception as exc: raise ValueError(f'{sql} -> {translated}: {exc}') from exc
finally: db.close()
return result
def validate_context_sql(sql,backend,schema):
"""Compile against the supplied fictitious schema during offline evaluation/auditing."""
import sqlglot
db=sqlite3.connect(':memory:');db.create_function('YEAR',1,lambda x:0);db.create_function('MONTH',1,lambda x:0)
try:
for ddl in schema:db.execute(ddl)
db.execute('PRAGMA query_only=ON')
db.execute(sqlite_sql(sqlglot.parse_one(sql,read='postgres' if backend=='supabase' else 'mysql')))
finally:db.close()
def sqlite_sql(tree):
from sqlglot import exp
def rewrite(node):
if isinstance(node,exp.Extract):
unit=str(node.this).upper()
if unit in ('YEAR','MONTH'):
return exp.Cast(this=exp.Anonymous(this='STRFTIME',expressions=[
exp.Literal.string('%Y' if unit=='YEAR' else '%m'),node.expression.copy()]),
to=exp.DataType.build('INTEGER'))
return node
return tree.copy().transform(rewrite).sql(dialect='sqlite')
def load_templates(path):
templates={}
for line in Path(path).read_text().splitlines():
row=json.loads(line); templates[row['op']]=row['templates']
missing=set(RECIPES)-set(templates)
if missing: raise ValueError(f'Missing teacher recipes: {sorted(missing)}')
return templates
def choose_template(templates,op,lang,split,rng):
pool=templates[op][lang]
if len(pool)<4: raise ValueError(f'Need four templates for {op}/{lang}, have {len(pool)}')
indices=list(range(len(pool)-2)) if split=='train' else [len(pool)-(2 if split=='validation' else 1)]
idx=rng.choice(indices)
return pool[idx],idx
def build_record(op,domain,index,backend,split,templates,rng):
s=schema_for(domain,index)
number=rng.choice([0,1,2,3,5,10,20,50,100,200,500,1000,1500,5000,10000])
s.update(number=number,upper=number+rng.choice([10,50,100,500]),limit=rng.choice([1,3,5,10,20]),
value=rng.choice(VALUES),other=rng.choice(VALUES),date=f'202{rng.randrange(2,7)}-{rng.randrange(1,13):02d}-15',
year=rng.randrange(2021,2027),month=rng.randrange(1,13),city=rng.choice(['Delhi','Mumbai','Lucknow','Jaipur']),
unit=rng.choice(['celsius','fahrenheit']),query=rng.choice(['row level security','SQL joins','API authentication','database backups','indexes']),
path=rng.choice(['/docs/README.md','/reports/sales.csv','/notes/setup.txt','/config/app.json']),
ticket_id='TKT-'+str(rng.randrange(100,10000)),count=rng.randrange(0,150))
project_id='demo_'+domain
tools,key,project,scoped=make_tools(backend,rng,split,project_id)
project_args={k:project_id for k in project}
context={'backend':backend,'project_id':project_id,'schema':ddls(s), 'tools':tools,
'policy':'Read-only database access. Use provided tools. Ask when required information is missing.'}
if op in SQL_OPS:
sql=gold_sql(op,s,backend)
validate_sql(sql,backend,s)
target={'action':'call','name':tools[0]['name'],'arguments':{**project_args,key:sql}}
elif op in ('list_tables','missing_schema'):
if op=='missing_schema': context['schema']=[]
target={'action':'call','name':tools[1]['name'],'arguments':{**project_args,'schemas':['public']}}
elif op in ('describe','sql_error'):
if op=='sql_error':
context['history']=[{'role':'tool','isError':True,'content':'Unknown column in previous query. Inspect the schema.'}]
context['schema']=[]
target=({'action':'call','name':tools[2]['name'],'arguments':{**project_args,'table':domain}}
if backend=='mysql' else
{'action':'call','name':tools[1]['name'],'arguments':{**project_args,'schemas':['public']}})
elif op in ('ambiguous','missing_value','write'):
target={'action':'clarify','question':{'ambiguous':'What does best mean: which column and order?',
'missing_value':f"Which value should {s['category']} equal?",'write':'This connection is read-only. Would you like a SELECT query instead?'}[op]}
elif op=='final':
context['history']=[{'role':'tool','content':{'row_count':s['count']}}]
target={'action':'answer','text':f"The tool returned {s['count']} rows."}
else:
definitions={
'weather':('get_weather','Get current weather for a city.',{'city':{'type':'string'},'unit':{'type':'string','enum':['celsius','fahrenheit']}}, {'city':s['city'],'unit':s['unit']}),
'search':('search_docs','Search documentation for a query.',{'query':{'type':'string'}},{'query':s['query']}),
'read_file':('read_file','Read text from a file path.',{'path':{'type':'string'}},{'path':s['path']}),
'ticket':('get_ticket','Retrieve a support ticket by its ID.',{'ticket_id':{'type':'string'}},{'ticket_id':s['ticket_id']}),
}
generic=[]
for kind,(name,desc,properties,args) in definitions.items():
prefix=rng.choice(['','tools.','app.']) if split=='train' else ('external.' if split=='test' else 'service.')
generic.append(tool(prefix+name,desc,properties))
if kind==op: target={'action':'call','name':prefix+name,'arguments':args}
tools.extend(generic)
context['schema']=[]
rng.shuffle(tools)
semantic_id=hashlib.sha256(compact([split,backend,s,op,target]).encode()).hexdigest()[:20]
records=[]
for lang in LANGUAGES:
template,template_index=choose_template(templates,op,lang,split,rng)
question=template.format(**s)
answer=dict(target)
if lang=='hi' and answer['action']=='clarify':
answer['question']={'ambiguous':'सबसे अच्छा से आपका क्या मतलब है? कौन सा कॉलम और क्रम?',
'missing_value':f"{s['category']} का कौन सा मान चाहिए?",'write':'यह कनेक्शन केवल पढ़ने के लिए है। क्या आपको SELECT क्वेरी चाहिए?'}[op]
elif lang=='hinglish' and answer['action']=='clarify':
answer['question']={'ambiguous':'Best ka matlab kya hai? Kaunsa column aur order?',
'missing_value':f"{s['category']} ki kaunsi value chahiye?",'write':'Yeh connection read-only hai. Kya SELECT query chahiye?'}[op]
elif answer['action']=='answer':
if lang=='hi': answer['text']=f"टूल ने {s['count']} पंक्तियाँ लौटाईं।"
elif lang=='hinglish': answer['text']=f"Tool ne {s['count']} rows return ki."
records.append({'id':semantic_id+'_'+lang,'scenario_id':semantic_id,'split':split,'domain':domain,
'language':lang,'backend':backend,'operation':op,'slots':s,'context':context,
'question':question,'target':answer,'prompt':serialize(context,question),
'response':compact(answer),'template_index':template_index,
'provenance':'programmatic semantics + Qwen3.8-27B-FP8 language template'})
return records
def main():
p=argparse.ArgumentParser(); p.add_argument('--templates',required=True)
p.add_argument('--out',default='data/tinyquery'); p.add_argument('--scenarios',type=int,default=18000)
args=p.parse_args(); out=Path(args.out); out.mkdir(parents=True,exist_ok=True)
templates=load_templates(args.templates); stats={}; global_prompts={}
for split,count in [('train',args.scenarios),('validation',300),('test',300)]:
rng=random.Random({'train':42,'validation':43,'test':44}[split]); seen=set(); counts={}
with (out/(split+'.jsonl')).open('w') as stream:
for i in range(count):
op=list(RECIPES)[i%len(RECIPES)]; domain=rng.choice(DOMAINS[split])
for record in build_record(op,domain,i,rng.choice(['mysql','supabase']),split,templates,rng):
key=hashlib.sha256(record['prompt'].encode()).hexdigest()
if key in seen: continue
if key in global_prompts and global_prompts[key]!=split: raise ValueError('Cross-split prompt leakage')
seen.add(key); global_prompts[key]=split
stream.write(json.dumps(record,ensure_ascii=False)+'\n')
counts[op]=counts.get(op,0)+1
if i%1000==0: print(split,i,flush=True)
stats[split]={'examples':len(seen),'operations':counts,'domains':DOMAINS[split]}
stats['sql_validation_method']='Dialect parsing plus transpiled SQLite execution on two fixtures; not native MySQL/Postgres execution.'
(out/'dataset-stats.json').write_text(json.dumps(stats,indent=2))
print(json.dumps(stats),flush=True)
if __name__=='__main__': main()