TinyQuery-140M / tinyquery /verify_concrete.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
3.87 kB
"""Round-trip direct paraphrases to SQL without exposing reference SQL to the verifier."""
import argparse
import concurrent.futures
import json
import time
import urllib.request
from pathlib import Path
from tinyquery.data import compact,validate_sql
def normalized(rows,ordered):
return rows if ordered else sorted(rows,key=lambda row:repr(row))
def verify(item,base):
records=item['accepted']
if not records: return {'scenario_id':item['scenario_id'],'accepted':[],'rejected':[]}
sample=records[0]
payload={'model':'Qwen/Qwen3.8-27B-FP8','messages':[{'role':'user','content':
'Translate each request to one read-only '+('PostgreSQL' if sample['backend']=='supabase' else 'MySQL')+
' SELECT query using only this schema. Return a JSON object mapping each language key to its SQL string. '
'If a request is ambiguous or impossible, use an empty string. Do not add constraints. '+
compact({'schema':sample['context']['schema'],'requests':{r['language']:r['question'] for r in records}})}],
'max_tokens':700,'temperature':0,'response_format':{'type':'json_object'},
'chat_template_kwargs':{'enable_thinking':False}}
request=urllib.request.Request(base+'/chat/completions',data=compact(payload).encode(),headers={'Content-Type':'application/json'})
with urllib.request.urlopen(request,timeout=180) as response: raw=json.load(response)
answers=json.loads(raw['choices'][0]['message']['content'])
gold=next(v for k,v in sample['target']['arguments'].items() if k in ['query','sql'])
reference=validate_sql(gold,sample['backend'],sample['slots'])
ordered=sample['operation'] in ['top','bottom','sort_asc','sort_desc']
accepted=[]; rejected=[]
for record in records:
try:
sql=answers[record['language']]
results=validate_sql(sql,record['backend'],record['slots'])
if any(normalized(a,ordered)!=normalized(b,ordered) for a,b in zip(reference,results)):
raise ValueError('SQL results differ on fixtures')
record['roundtrip_sql']=sql
record['provenance']+='; teacher round-trip SQL equivalent on two generated SQLite fixtures'
accepted.append(record)
except Exception as exc:
rejected.append({'id':record['id'],'question':record['question'],'reason':str(exc),'sql':answers.get(record['language'])})
return {'scenario_id':item['scenario_id'],'accepted':accepted,'rejected':rejected,'verification_request':payload,'verification_response':raw}
def main():
p=argparse.ArgumentParser(); p.add_argument('--input',required=True); p.add_argument('--out',required=True)
p.add_argument('--workers',type=int,default=24); p.add_argument('--base',default='http://127.0.0.1:18000/v1')
args=p.parse_args(); output=Path(args.out); done=set()
if output.exists(): done={json.loads(l)['scenario_id'] for l in output.read_text().splitlines()}
items=[json.loads(l) for l in Path(args.input).read_text().splitlines()]
start=time.time(); completed=0; accepted=0; rejected=0
with output.open('a') as stream,concurrent.futures.ThreadPoolExecutor(args.workers) as pool:
jobs={pool.submit(verify,r,args.base):r['scenario_id'] for r in items if r['scenario_id'] not in done}
for future in concurrent.futures.as_completed(jobs):
completed+=1
try:
result=future.result(); accepted+=len(result['accepted']); rejected+=len(result['rejected'])
stream.write(compact(result)+'\n'); stream.flush()
except Exception as exc: print(compact({'error':str(exc),'scenario_id':jobs[future]}),flush=True)
if completed%25==0: print(compact({'completed':completed,'accepted':accepted,'rejected':rejected,'seconds':time.time()-start}),flush=True)
if __name__=='__main__': main()