TinyQuery-140M / tinyquery /baseline.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.34 kB
"""Transparent TF-IDF nearest-question baseline; copies the retrieved action unchanged."""
import argparse
from collections import Counter,defaultdict
import json
import math
from pathlib import Path
import re
import time
import numpy as np
from tinyquery.evaluate import score,aggregate
def bind_context(text,source,destination):
"""Bind by tool descriptions and DDL positions, using no expected answer or slots."""
import sqlglot
from sqlglot import exp
action=json.loads(text)
if action.get('action')!='call':return text
old=next(t for t in source['tools'] if t['name']==action['name'])
normalize=lambda s:s.replace('PostgreSQL','SQL').replace('MySQL','SQL')
candidates=[t for t in destination['tools'] if normalize(t['description'])==normalize(old['description'])]
if not candidates:return text
tool=candidates[0];props=tool['inputSchema']['properties'];args=dict(action['arguments'])
sqlkey=next((k for k in ['sql','query'] if isinstance(args.get(k),str) and args[k].startswith('SELECT ')),None)
if sqlkey:
targetkey=next((k for k in ['sql','query'] if k in props),sqlkey)
tables={};columns={}
for old_ddl,new_ddl in zip(source['schema'],destination['schema']):
a=sqlglot.parse_one(old_ddl).this;b=sqlglot.parse_one(new_ddl).this
tables[a.this.name]=b.this.name
for c,d in zip(a.expressions,b.expressions):
if isinstance(c,exp.ColumnDef) and isinstance(d,exp.ColumnDef):columns[c.name]=d.name
query=sqlglot.parse_one(args.pop(sqlkey),read='postgres' if source['backend']=='supabase' else 'mysql')
def rename(node):
if isinstance(node,exp.Table) and node.name in tables:node.set('this',exp.to_identifier(tables[node.name]))
if isinstance(node,exp.Column):
if node.name in columns:node.set('this',exp.to_identifier(columns[node.name]))
if node.table in tables:node.set('table',exp.to_identifier(tables[node.table]))
return node
args[targetkey]=query.transform(rename).sql(dialect='postgres' if destination['backend']=='supabase' else 'mysql')+';'
args={k:v for k,v in args.items() if k in props}
if 'project_id' in props:args['project_id']=destination['project_id']
return json.dumps({'action':'call','name':tool['name'],'arguments':args},ensure_ascii=False,separators=(',',':'))
def words(text):
return re.findall(r'[^\W_]+|_',text.lower(),re.UNICODE)
def main():
p=argparse.ArgumentParser(); p.add_argument('--train',required=True); p.add_argument('--data',required=True)
p.add_argument('--out',required=True);p.add_argument('--bind-context',action='store_true')
args=p.parse_args(); start=time.time()
training=[]; seen=set(); counts=[]; df=Counter()
for line in Path(args.train).open():
row=json.loads(line)
if row['question'] in seen: continue
seen.add(row['question']); terms=Counter(words(row['question']))
training.append((row['id'],row['response'],row['context'] if args.bind_context else None)); counts.append(terms); df.update(terms.keys())
n=len(training); idf={word:math.log((1+n)/(1+freq))+1 for word,freq in df.items()}
postings=defaultdict(list)
for i,terms in enumerate(counts):
weighted={word:(1+math.log(freq))*idf[word] for word,freq in terms.items()}
norm=math.sqrt(sum(v*v for v in weighted.values())) or 1
for word,value in weighted.items(): postings[word].append((i,value/norm))
postings={word:(np.array([i for i,_ in pairs]),np.array([v for _,v in pairs],dtype=np.float32))
for word,pairs in postings.items()}
results=[]; groups=defaultdict(list); out=Path(args.out); out.parent.mkdir(parents=True,exist_ok=True)
with out.open('w') as stream:
for line in Path(args.data).open():
row=json.loads(line); similarities=np.zeros(n,dtype=np.float32)
for word,freq in Counter(words(row['question'])).items():
if word in postings:
indices,weights=postings[word]; similarities[indices]+=(1+math.log(freq))*idf[word]*weights
index=int(similarities.argmax()); source,text,context=training[index]
if args.bind_context:
try:text=bind_context(text,context,row['context'])
except (ValueError,KeyError,StopIteration,AttributeError):pass
metrics=score(row,text)
results.append(metrics)
for field in ['language','backend','operation']: groups[field+':'+row[field]].append(metrics)
stream.write(json.dumps({'id':row['id'],'retrieved_id':source,'output':text,'metrics':metrics},ensure_ascii=False)+'\n')
report={'baseline':'TF-IDF nearest training question; copy action unchanged, no tool/schema adaptation',
'training_questions':n,'examples':len(results),'seconds':time.time()-start,'metrics':aggregate(results),
'groups':{key:aggregate(value) for key,value in groups.items()}}
if args.bind_context:report['baseline']='TF-IDF nearest question plus tool-description and DDL-position binding; no expected answer or gold slots used'
out.with_suffix('.summary.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report['metrics']))
if __name__=='__main__': main()