TinyQuery-140M / tinyquery /evaluate.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
9.36 kB
"""Raw greedy tool-calling evaluation; SQL equivalence is explicitly SQLite-based."""
import argparse
from collections import Counter,defaultdict
import contextlib
import hashlib
import json
import time
from pathlib import Path
import torch
import jsonschema
from tokenizers import Tokenizer
from tinyquery.model import TinyQuery,Config
from tinyquery.data import validate_sql,validate_context_sql,SQL_OPS
def parse_action(text):
def unique(pairs):
result={}
for key,value in pairs:
if key in result:raise ValueError('Duplicate JSON key: '+key)
result[key]=value
return result
return json.loads(text,object_pairs_hook=unique)
def check_action(action,context):
if not isinstance(action,dict): raise ValueError('Action must be an object')
kind=action.get('action')
if kind=='call':
if set(action)!={'action','name','arguments'}: raise ValueError('Unexpected action keys')
tool=next((t for t in context['tools'] if t['name']==action.get('name')),None)
if tool is None: raise ValueError('Unknown tool')
jsonschema.validate(action.get('arguments'),tool['inputSchema'])
elif kind in ('clarify','answer'):
key='question' if kind=='clarify' else 'text'
if set(action)!={'action',key} or not isinstance(action[key],str) or not action[key].strip():
raise ValueError('Invalid textual action')
else: raise ValueError('Unknown action')
return action
def mcp_request(action,context,request_id=1):
check_action(action,context)
if action['action']!='call': raise ValueError('Only tool actions map to tools/call')
# Scope is an inference input, not a string the model may freely substitute.
if 'project_id' in action['arguments'] and action['arguments']['project_id']!=context.get('project_id'):
raise ValueError('Predicted project_id differs from the supplied project scope')
tool=next(t for t in context['tools'] if t['name']==action['name'])
description=tool.get('description','')
for key,value in action['arguments'].items():
if key=='sql' or (key=='query' and 'SELECT query' in description):
import sqlglot
from sqlglot import exp
trees=sqlglot.parse(value,read='mysql' if context['backend']=='mysql' else 'postgres')
if len(trees)!=1 or not isinstance(trees[0],exp.Select):raise ValueError('Expected one read-only SELECT')
tree=trees[0]
if tree.find(exp.Into) or any(isinstance(n,(exp.DML,exp.DDL)) for n in tree.walk()):
raise ValueError('Data-changing SQL is not supported')
known={sqlglot.parse_one(ddl).this.this.name for ddl in context.get('schema',[])}
if any(table.name not in known for table in tree.find_all(exp.Table)):
raise ValueError('SQL references a table absent from the supplied schema')
return {'jsonrpc':'2.0','id':request_id,'method':'tools/call',
'params':{'name':action['name'],'arguments':action['arguments']}}
def score(row,text):
result={'json_valid':False,'schema_valid':False,'action_correct':False,'tool_correct':False,
'arguments_exact':False,'sql_equivalent':None,'success':False}
if row['target']['action']!='call':
result['tool_correct']=None; result['arguments_exact']=None
if row['operation'] in SQL_OPS: result['sql_equivalent']=False
try:
parsed=parse_action(text); result['json_valid']=True
if isinstance(parsed,dict): result['action_correct']=parsed.get('action')==row['target']['action']
check_action(parsed,row['context']); result['schema_valid']=True
gold=row['target']; result['action_correct']=parsed['action']==gold['action']
if not result['action_correct']: return result
if gold['action']!='call':
key='question' if gold['action']=='clarify' else 'text'
# Action accuracy and wording exact-match are separate; generic clarification is not semantic proof.
result['text_exact']=parsed[key]==gold[key]
result['success']=result['text_exact']
return result
result['tool_correct']=parsed['name']==gold['name']
result['arguments_exact']=parsed['arguments']==gold['arguments']
sqlkey=next((k for k in ('sql','query') if k in gold['arguments'] and str(gold['arguments'][k]).startswith('SELECT ')),None)
if sqlkey:
# Backend/project arguments must also match; equal SQL against a wrong project is not a pass.
other_correct={k:v for k,v in parsed['arguments'].items() if k!=sqlkey}=={k:v for k,v in gold['arguments'].items() if k!=sqlkey}
validate_context_sql(parsed['arguments'][sqlkey],row['backend'],row['context']['schema'])
actual=validate_sql(parsed['arguments'][sqlkey],row['backend'],row['slots'])
expected=validate_sql(gold['arguments'][sqlkey],row['backend'],row['slots'])
ordered=row['operation'] in ['sort_asc','sort_desc','top','bottom']
normalize=lambda a:a if ordered else sorted(a,key=repr)
equivalent=all(normalize(a)==normalize(b) for a,b in zip(actual,expected))
result['sql_equivalent']=equivalent
result['success']=result['tool_correct'] and other_correct and equivalent
else: result['success']=result['tool_correct'] and result['arguments_exact']
except Exception as exc: result['error']=str(exc)[:350]
return result
def load_model(checkpoint,device):
path=Path(checkpoint)
if path.suffix=='.safetensors':
from safetensors.torch import load
c=Config(**json.loads((path.parent/'config.json').read_text()))
raw=path.read_bytes(); state=load(raw)
model=TinyQuery(c); model.load_state_dict(state)
model.checkpoint_sha256=hashlib.sha256(raw).hexdigest();del raw,state
else:
state=torch.load(path,map_location='cpu',weights_only=False)
model=TinyQuery(Config(**state['config'])); model.load_state_dict(state['model']); del state
model=model.to(device)
if device=='cuda': model=model.to(torch.bfloat16)
return model.eval()
def aggregate(results):
metrics={}
for key in ['json_valid','schema_valid','action_correct','tool_correct','arguments_exact','sql_equivalent','success']:
entries=[r[key] for r in results if r.get(key) is not None]
metrics[key]={'correct':sum(entries),'total':len(entries),'rate':sum(entries)/max(1,len(entries))}
return metrics
def main():
p=argparse.ArgumentParser(); p.add_argument('--checkpoint',required=True); p.add_argument('--tokenizer',required=True)
p.add_argument('--data',required=True); p.add_argument('--out',required=True); p.add_argument('--limit',type=int,default=0)
p.add_argument('--batch',type=int,default=16); p.add_argument('--tokens',type=int,default=160)
args=p.parse_args(); device='cuda' if torch.cuda.is_available() else ('mps' if torch.backends.mps.is_available() else 'cpu')
torch.set_num_threads(8); model=load_model(args.checkpoint,device); tokenizer=Tokenizer.from_file(args.tokenizer)
rows=[json.loads(l) for l in Path(args.data).read_text().splitlines()]
if args.limit:
import random
random.Random(777).shuffle(rows); rows=rows[:args.limit]
# Sort only for efficient padding, never alter which examples are evaluated.
encoded=[(r,tokenizer.encode(r['prompt']).ids) for r in rows]; encoded.sort(key=lambda item:len(item[1]))
out=Path(args.out); out.parent.mkdir(parents=True,exist_ok=True); results=[]; start=time.time()
groups=defaultdict(list)
with out.open('w') as stream:
for i in range(0,len(encoded),args.batch):
batch=encoded[i:i+args.batch]
with torch.inference_mode():
outputs=model.generate_batch([ids for _,ids in batch],tokenizer.token_to_id('<|end|>'),max_new_tokens=args.tokens)
for (row,_),ids in zip(batch,outputs):
text=tokenizer.decode(ids,skip_special_tokens=True)
metrics=score(row,text); results.append(metrics)
for field in ['language','backend','operation']: groups[field+':'+row[field]].append(metrics)
entry={'id':row['id'],'language':row['language'],'backend':row['backend'],'operation':row['operation'],
'question':row['question'],'expected':row['target'],'output':text,'metrics':metrics}
stream.write(json.dumps(entry,ensure_ascii=False)+'\n')
stream.flush()
print(json.dumps({'evaluated':len(results),'seconds':time.time()-start,'success':sum(r['success'] for r in results)/len(results)}),flush=True)
summary={'examples':len(results),'seconds':time.time()-start,'decoding':'raw greedy, no repair, no teacher fallback',
'checkpoint_sha256':getattr(model,'checkpoint_sha256',None),
'sql_metric':'Compile against the supplied schema, then compare results on two generated SQLite fixtures after dialect adaptation; not native MySQL/PostgreSQL execution.',
'metrics':aggregate(results),'groups':{k:aggregate(v) for k,v in groups.items()}}
out.with_suffix('.summary.json').write_text(json.dumps(summary,indent=2)); print(json.dumps(summary),flush=True)
if __name__=='__main__': main()