Files
llm-speed-tester/tester.py
T

248 lines
11 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""速度测试执行器:校准 -> 采样 -> 汇总,全程写日志与指标入库"""
import json
import statistics
import threading
import time
import uuid
import database as db
import llm_providers as lp
from llm_providers import ProviderError, StopRequested
class TestRunner(threading.Thread):
def __init__(self, test_id, cfg, gen):
super().__init__(daemon=True)
self.test_id = test_id
self.cfg = cfg
self.gen = gen
self.cancel_flag = False
self.start_wall = time.time()
self.ratio = None
self.samples = []
self.last_error = None
def request_cancel(self):
self.cancel_flag = True
def should_stop(self):
return self.cancel_flag
def log(self, level, msg):
db.add_log(self.test_id, level, msg)
# ───────────────────────── 主流程 ─────────────────────────
def run(self):
try:
self._run()
except StopRequested:
self.log("WARN", "用户请求停止测试")
db.update_status(self.test_id, "canceled",
summary=self._make_summary(), error="用户取消")
except Exception as e:
self.log("ERROR", "测试异常终止: %s" % e)
db.update_status(self.test_id, "error",
summary=self._make_summary(), error=str(e))
def _run(self):
provider = self.cfg.get("provider", "openai")
model = self.cfg.get("model", "")
gen = self.gen
# 上下文长度列表(支持手动自定义,默认 512/2048/4096/8192/16384/32768/65536/131072
raw_lengths = gen.get("context_lengths") or []
if not raw_lengths:
# 兼容旧版单值配置
raw_lengths = [int(gen.get("prompt_tokens", 2048))]
lengths = sorted(set(int(x) for x in raw_lengths if int(x) >= 16)) or [2048]
n = max(1, int(gen.get("samples", 2))) # 每个长度采样次数
max_tokens = max(1, int(gen.get("max_tokens", 128))) # 解码输出长度
avoid_cache = bool(gen.get("avoid_cache"))
warmup = bool(gen.get("warmup", True)) # 测试前空转预热
self.log("INFO", "═══ 开始速度测试 ═══")
name = gen.get("name") or self.cfg.get("name") or ""
if name:
self.log("INFO", "测试名称(主题): %s" % name)
self.log("INFO", "提供商: %s | 模型: %s" % (lp.PROVIDER_LABELS.get(provider, provider), model))
self.log("INFO", "上下文长度: %s tokens | 生成长度: %d tokens | 每个长度采样: %d 次 | 预热: %s | 避免缓存: %s"
% (" / ".join(str(x) for x in lengths), max_tokens, n,
"开" if warmup else "关", "开" if avoid_cache else "关"))
ratio = self._calibrate()
self.ratio = ratio
self.log("INFO", "校准完成: %.3f tok/字符(%.2f 字符/token" % (ratio, 1.0 / ratio))
for L in lengths:
if self.should_stop():
raise StopRequested()
base_prompt = self._build_prompt(L, ratio)
self.log("INFO", "▸▸ 上下文长度 %d tokens(基准提示词构造完成)" % L)
if warmup:
self._warmup(base_prompt)
for i in range(1, n + 1):
if self.should_stop():
raise StopRequested()
prompt = self._finalize_prompt(base_prompt)
self.log("INFO", "── [%d tok] 采样 %d/%d 开始 ──" % (L, i, n))
try:
m = lp.call_stream(
self.cfg, prompt,
{"max_tokens": max_tokens, "avoid_cache": avoid_cache},
log=lambda lv, msg: self.log(lv, msg),
should_stop=self.should_stop)
m["run_index"] = i
m["context_length"] = L
self.samples.append({"run_index": i, "context_length": L, "ok": True, "metrics": m})
db.add_run(self.test_id, i, m, context_length=L)
self.log("METRIC", self._fmt_metric(L, i, n, m))
except StopRequested:
raise
except ProviderError as e:
# 单次采样失败:记录并继续后续采样,不让整个测试中断
self.last_error = str(e)
self.log("ERROR", "[%d tok] 采样 %d/%d 失败: %s" % (L, i, n, e))
self.samples.append({"run_index": i, "context_length": L, "ok": False, "error": str(e)})
db.add_run(self.test_id, i, {}, str(e), context_length=L)
summary = self._make_summary()
ok_count = summary.get("samples_ok") or 0
fail_count = summary.get("samples_total", 0) - ok_count
if ok_count:
db.update_status(self.test_id, "done", summary=summary,
error=("%d 次采样失败:%s" % (fail_count, self.last_error)) if fail_count else "")
self.log("INFO", "═══ 测试完成 ═══")
if fail_count:
self.log("WARN", "共 %d 次采样失败(最后错误:%s" % (fail_count, self.last_error))
else:
db.update_status(self.test_id, "error", summary=summary,
error=self.last_error or "所有采样均失败")
self.log("ERROR", "所有采样均失败,测试标记为 error(最后错误:%s" % (self.last_error or "未知"))
return
self.log("INFO", "汇总: 平均首字 %.1f ms | 平均预填充 %.1f tok/s | 平均解码 %.1f tok/s"
% (summary.get("avg_ttft_ms") or 0,
summary.get("avg_prefill_speed") or 0,
summary.get("avg_decode_speed") or 0))
def _warmup(self, base_prompt):
"""空转预热:不计入任何速度统计,用于避免冷启动/首次请求偏慢影响采样"""
self.log("INFO", "预热(空转,不计速度)...")
try:
lp.call_stream(self.cfg, base_prompt,
{"max_tokens": 8, "avoid_cache": False},
log=lambda lv, msg: self.log(lv, msg),
should_stop=self.should_stop)
self.log("INFO", "预热完成(不纳入统计)")
except StopRequested:
raise
except Exception as e:
self.log("WARN", "预热失败(继续测试): %s" % e)
# ───────────────────────── 工具方法 ─────────────────────────
def _calibrate(self):
probe = ("The quick brown fox jumps over the lazy dog. 人工智能大模型推理速度基准语料,"
"用于测量提示词预填充与流式解码性能。\n") * 40
self.log("INFO", "正在校准 token/字符 比例(发送小探测请求)...")
try:
m = lp.call_stream(self.cfg, probe,
{"max_tokens": 8, "avoid_cache": False},
log=lambda lv, msg: self.log(lv, msg),
should_stop=self.should_stop)
pt = m.get("prompt_tokens") or 0
if pt and len(probe):
ratio = pt / len(probe)
self.log("INFO", "探测提示词 %d tokens / %d 字符 = %.3f tok/字符"
% (pt, len(probe), ratio))
return max(ratio, 0.001)
except StopRequested:
raise
except Exception as e:
self.log("WARN", "校准失败(%s),使用默认估算 0.55 tok/字符" % e)
return 0.55
def _build_prompt(self, target_tokens, ratio):
seg = ("基准语料:The quick brown fox jumps over the lazy dog. "
"人工智能大模型推理性能测试文本,用于测量提示词预填充速度、首字延迟与流式解码吞吐。\n")
target_chars = max(64, int(target_tokens / ratio))
repeats = max(1, target_chars // len(seg))
return seg * repeats
def _finalize_prompt(self, base):
if self.gen.get("avoid_cache"):
return "[cache-bust %s]\n%s" % (uuid.uuid4().hex, base)
return base
def _fmt_metric(self, L, i, n, m):
return ("[%d tok] 采样 %d/%d 完成 | 提示词 %d tok | 缓存 %d tok | 首字 %s ms | 预填充 %s tok/s"
" | 输出 %d tok | 解码 %s tok/s | 总耗时 %s ms"
% (L, i, n, m.get("prompt_tokens") or 0, m.get("cached_tokens") or 0,
m.get("ttft_ms"), m.get("prefill_speed"), m.get("output_tokens") or 0,
m.get("decode_speed"), m.get("total_ms")))
def _make_summary(self):
ok = [s for s in self.samples if s.get("ok")]
base = {
"provider": self.cfg.get("provider"),
"model": self.cfg.get("model"),
"gen": self.gen,
"samples_total": len(self.samples),
"samples_ok": len(ok),
"calibration_chars_per_token": round(1 / self.ratio, 2) if self.ratio else None,
}
if not ok:
return base
def avg(ms, k):
vals = [m[k] for m in ms if m.get(k) is not None]
return round(statistics.mean(vals), 1) if vals else None
# 按上下文长度分组汇总
by_length = {}
for L in sorted(set(s["context_length"] for s in ok)):
group = [s["metrics"] for s in ok if s["context_length"] == L]
by_length[L] = {
"samples_total": sum(1 for s in self.samples if s["context_length"] == L),
"samples_ok": len(group),
"avg_ttft_ms": avg(group, "ttft_ms"),
"avg_prefill_speed": avg(group, "prefill_speed"),
"avg_decode_speed": avg(group, "decode_speed"),
"avg_prompt_tokens": avg(group, "prompt_tokens"),
"avg_output_tokens": avg(group, "output_tokens"),
"avg_total_ms": avg(group, "total_ms"),
}
okm = [s["metrics"] for s in ok]
def mn(k):
vals = [m[k] for m in okm if m.get(k) is not None]
return round(min(vals), 1) if vals else None
def mx(k):
vals = [m[k] for m in okm if m.get(k) is not None]
return round(max(vals), 1) if vals else None
summary = dict(base)
summary.update({
"by_length": by_length,
"avg_ttft_ms": avg(okm, "ttft_ms"),
"min_ttft_ms": mn("ttft_ms"),
"max_ttft_ms": mx("ttft_ms"),
"avg_prefill_speed": avg(okm, "prefill_speed"),
"min_prefill_speed": mn("prefill_speed"),
"max_prefill_speed": mx("prefill_speed"),
"avg_decode_speed": avg(okm, "decode_speed"),
"min_decode_speed": mn("decode_speed"),
"max_decode_speed": mx("decode_speed"),
"avg_prompt_tokens": avg(okm, "prompt_tokens"),
"avg_output_tokens": avg(okm, "output_tokens"),
"avg_cached_tokens": avg(okm, "cached_tokens"),
"avg_total_ms": avg(okm, "total_ms"),
"min_total_ms": mn("total_ms"),
"max_total_ms": mx("total_ms"),
"best_ttft_ms": mn("ttft_ms"),
})
return summary