# -*- 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/8192/32768/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