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.71 kB
"""Stream the student's raw JSON response from a checkpoint; optionally render an MCP request."""
import argparse
import json
import sys
import time
from pathlib import Path
import torch
from tokenizers import Tokenizer,decoders
from tinyquery.data import serialize,make_tools
from tinyquery.evaluate import load_model,check_action,mcp_request,parse_action
def default_context(backend):
import random
tools,_,_,_=make_tools(backend,random.Random(4),'train','demo')
return {'backend':backend,'project_id':'demo',
'schema':['CREATE TABLE customers (id INTEGER PRIMARY KEY, name TEXT, city TEXT, amount REAL);'],
'tools':tools,'policy':'Read-only database access. Use provided tools. Ask when required information is missing.'}
def stream_response(model,tokenizer,prompt,max_tokens=160,stats=None):
ids=tokenizer.encode(prompt).ids
if len(ids)+max_tokens>model.config.context: raise ValueError('Prompt plus output budget exceeds model context')
device=next(model.parameters()).device; tokens=torch.tensor([ids],device=device)
past=None; decoder=decoders.DecodeStream(skip_special_tokens=True)
generated=[]; emitted='';start=time.perf_counter();first_token=None
with torch.inference_mode():
for _ in range(max_tokens):
logits,past,_=model(tokens,past=past,use_cache=True,last_only=True)
token=int(logits[0,-1].argmax()); generated.append(token)
if first_token is None:first_token=time.perf_counter()-start
text=decoder.step(tokenizer,token)
if text: emitted+=text; yield text
if token==tokenizer.token_to_id('<|end|>'): break
tokens=torch.tensor([[token]],device=device)
full=tokenizer.decode(generated,skip_special_tokens=True)
if full.startswith(emitted) and len(full)>len(emitted): yield full[len(emitted):]
if stats is not None:
seconds=time.perf_counter()-start
stats.update(generated_tokens=len(generated),seconds=seconds,tokens_per_second=len(generated)/seconds,
first_token_seconds=first_token,device=str(device),includes_model_loading=False)
def main():
p=argparse.ArgumentParser(); p.add_argument('question'); p.add_argument('--checkpoint',required=True)
p.add_argument('--tokenizer'); p.add_argument('--context',help='JSON file with backend, schema and MCP-style tools')
p.add_argument('--backend',choices=['mysql','supabase'],default='supabase')
p.add_argument('--tokens',type=int,default=160); p.add_argument('--mcp',action='store_true')
p.add_argument('--stats',action='store_true',help='Print measured generation speed to stderr after streaming')
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(4)
tokenizer=Tokenizer.from_file(args.tokenizer or str(Path(args.checkpoint).parent/'tokenizer.json'))
model=load_model(args.checkpoint,device)
context=json.loads(Path(args.context).read_text()) if args.context else default_context(args.backend)
answer='';stats={}
for text in stream_response(model,tokenizer,serialize(context,args.question),args.tokens,stats):
print(text,end='',flush=True); answer+=text
print(flush=True)
if args.stats:print(json.dumps(stats),file=sys.stderr)
try:
action=check_action(parse_action(answer),context)
if args.mcp and action['action']=='call': print(json.dumps(mcp_request(action,context),ensure_ascii=False,indent=2))
except Exception as exc:
print('Output validation failed:',str(exc),file=sys.stderr)
raise SystemExit(1)
if __name__=='__main__': main()