Download tinyquery/teacher_bench.py from karmx/TinyQuery-140M: direct link, hf CLI and curl.
- Browser
- Download file 2.24 kB
-
https://huggingface.co/karmx/TinyQuery-140M/resolve/main/tinyquery/teacher_bench.py
- Command line
-
hf download hf://karmx/TinyQuery-140M/tinyquery/teacher_bench.py
-
curl -L -o teacher_bench.py https://huggingface.co/karmx/TinyQuery-140M/resolve/main/tinyquery/teacher_bench.py
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() | |