Lottolabs's picture
Upload verified mixed BFP4/BFP8 checkpoint with MTP and evaluation evidence
12f320c verified
Raw History Blame Contribute Delete
6.25 kB
"""Exercise native serving behavior, near-limit prefill and greedy decode over HTTP.
Local diagnostics, not an official task benchmark or Unsloth Divergence-300.
"""
import argparse
import json
import statistics
import time
import urllib.request
from pathlib import Path
from score_behavior import passed
def payload(model, messages, max_tokens):
return {"model": model, "messages": messages, "temperature": 0, "top_p": 1,
"seed": 9472, "max_tokens": max_tokens,
"chat_template_kwargs": {"enable_thinking": False}}
def send(base_url, body):
return urllib.request.urlopen(urllib.request.Request(
base_url.rstrip('/') + '/v1/chat/completions',
data=json.dumps(body).encode(), headers={"Content-Type": "application/json"}), timeout=600)
def check_case(base_url, model, case, max_tokens):
messages = case.get('messages', [{"role": "user", "content": case['prompt']}])
body = payload(model, messages, max_tokens)
if case.get('tools'):
body.update(tools=case['tools'], tool_choice='auto')
started = time.monotonic()
with send(base_url, body) as response:
result = json.load(response)
message = result['choices'][0]['message']
text = message.get('content') or ''
if case['check'] == 'weather_tool':
calls = message.get('tool_calls') or []
ok = len(calls) == 1 and calls[0]['function']['name'] == 'get_weather'
if ok:
arguments = json.loads(calls[0]['function']['arguments'])
ok = arguments.get('city', '').casefold() == str(case['expected']).casefold()
else:
ok = passed(case, text)
return {'id': case['id'], 'pass': bool(ok), 'seconds': time.monotonic() - started,
'response': result}
def benchmark(base_url, model):
body = payload(model, [{"role": "user", "content":
"Write a detailed technical explanation of why batch-one language-model decoding "
"is often memory-bandwidth limited. Explain weight traffic, activation traffic, "
"kernel launch overhead, and how these differ from prompt prefill. Use several paragraphs."}], 128)
body.update(stream=True, stream_options={'include_usage': True})
started = time.monotonic()
first_token, finished, usage = None, None, None
pieces = []
with send(base_url, body) as response:
for raw in response:
line = raw.decode().strip()
if not line.startswith('data: '):
continue
data = line[6:]
if data == '[DONE]':
break
event = json.loads(data)
if event.get('usage'):
usage = event['usage']
for choice in event.get('choices', []):
delta = choice.get('delta', {})
text = (delta.get('reasoning_content') or '') + (delta.get('content') or '')
if text:
if first_token is None:
first_token = time.monotonic()
pieces.append(text)
if choice.get('finish_reason') is not None:
finished = time.monotonic()
if first_token is None or finished is None or not usage or usage['completion_tokens'] < 2:
raise ValueError('Missing streamed timing or endpoint token-usage evidence')
return {'text': ''.join(pieces), 'usage': usage,
'ttft_ms': (first_token - started) * 1000,
'decode_tokens_per_second': (usage['completion_tokens'] - 1) / (finished - first_token),
'decode_seconds': finished - first_token,
'timing_scope': 'Client-observed first nonempty delta through finish event; excludes first output token'}
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument('--base-url', default='http://127.0.0.1:8001')
parser.add_argument('--model', default='Qwen/Qwen3.5-9B')
parser.add_argument('--cases', type=Path, required=True)
parser.add_argument('--long-case', type=Path, required=True)
parser.add_argument('--output', type=Path, required=True)
args = parser.parse_args()
report = {'status': 'running', 'scope': __doc__, 'base_url': args.base_url,
'model': args.model, 'behavior': [], 'long_context': [], 'decode_runs': []}
try:
report['initial_decode'] = benchmark(args.base_url, args.model)
for case in map(json.loads, args.cases.read_text().splitlines()):
result = check_case(args.base_url, args.model, case, 64)
report['behavior'].append(result)
print(json.dumps({'case': result['id'], 'pass': result['pass']}), flush=True)
for case in map(json.loads, args.long_case.read_text().splitlines()):
result = check_case(args.base_url, args.model, case, 16)
result['prepared_prompt_tokens'] = len(case['token_ids'])
report['long_context'].append(result)
print(json.dumps({'case': result['id'], 'pass': result['pass']}), flush=True)
for _ in range(3):
report['decode_runs'].append(benchmark(args.base_url, args.model))
report.update(status='complete', behavior_passes=sum(r['pass'] for r in report['behavior']),
long_context_passes=sum(r['pass'] for r in report['long_context']),
median_decode_tokens_per_second=statistics.median(r['decode_tokens_per_second'] for r in report['decode_runs']),
identical_greedy_completions=len({r['text'] for r in report['decode_runs']}) == 1)
report['request_isolation_passed'] = all(
run['text'] == report['initial_decode']['text'] for run in report['decode_runs'])
if not report['request_isolation_passed']:
raise RuntimeError('Greedy output changed after interleaved requests, including long-context prefill')
except BaseException as error:
report.update(status='failed', error=str(error))
raise
finally:
args.output.write_text(json.dumps(report, indent=2, ensure_ascii=False) + '\n')
print(json.dumps({k: v for k, v in report.items() if k not in ('behavior', 'long_context', 'decode_runs')}, indent=2), flush=True)
if __name__ == '__main__':
main()