TinyQuery-140M / tinyquery /teacher_bench.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
2.24 kB
"""Identical-request, warmed-cache throughput comparison at fixed concurrency."""
import argparse
import concurrent.futures
import json
import time
import urllib.request
from pathlib import Path
def main():
p=argparse.ArgumentParser(); p.add_argument('--input',required=True); p.add_argument('--out',required=True)
p.add_argument('--label',required=True); p.add_argument('--workers',type=int,default=32)
args=p.parse_args()
requests=[json.loads(l)['request'] for l in Path(args.input).read_text().splitlines()[:64]]
def call(payload):
payload=dict(payload); payload.update(temperature=.7,top_p=.8,top_k=20,presence_penalty=0,repetition_penalty=1.0)
request=urllib.request.Request('http://127.0.0.1:18000/v1/chat/completions',
data=json.dumps(payload).encode(),headers={'Content-Type':'application/json'})
with urllib.request.urlopen(request,timeout=180) as response: result=json.load(response)
try: json.loads(result['choices'][0]['message']['content']); valid=True
except Exception: valid=False
return {'tokens':result['usage']['completion_tokens'],'json_valid':valid,'finish_reason':result['choices'][0]['finish_reason']}
rounds=[]
with concurrent.futures.ThreadPoolExecutor(args.workers) as pool:
for repeat in range(3):
start=time.time(); outputs=list(pool.map(call,requests)); seconds=time.time()-start
entry={'round':repeat,'warmup':repeat==0,'seconds':seconds,'tokens':sum(o['tokens'] for o in outputs),
'json_valid':sum(o['json_valid'] for o in outputs),'requests':len(outputs),
'truncated':sum(o['finish_reason']=='length' for o in outputs)}
entry['output_tps']=entry['tokens']/seconds; rounds.append(entry); print(json.dumps(entry),flush=True)
with urllib.request.urlopen('http://127.0.0.1:18000/metrics',timeout=10) as response:
metrics=response.read().decode()
result={'label':args.label,'workers':args.workers,'rounds':rounds,
'mtp_metrics':[l for l in metrics.splitlines() if 'spec_decode' in l and not l.startswith('#')]}
Path(args.out).write_text(json.dumps(result,indent=2))
if __name__=='__main__': main()