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