"""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()