Initial commit of LLM speed test app

This commit is contained in:
2026-09-06 11:45:52 +08:00
commit 5544b5f9f9
38 changed files with 8587 additions and 0 deletions
+240
View File
@@ -0,0 +1,240 @@
"""测试引擎:编排测试用例、统计指标、保存结果并生成报告"""
import json
import logging
import os
import sys
import time
from datetime import datetime
from glob import glob
from . import config
from . import report as report_module
from .stats import mean, std, percentile, percentiles, min_val, max_val
from .ollama_client import OllamaClient, OllamaError
from .test_cases.case_01_generation import run_test as run_case_01
from .test_cases.case_02_simple_tool import run_test as run_case_02
from .test_cases.case_03_file_analysis import run_test as run_case_03
from .test_cases.case_04_parallel_tool import run_test as run_case_04
from .test_cases.case_05_long_context import run_test as run_case_05
from .test_cases.case_06_reasoning import run_test as run_case_06
from .test_cases.case_07_multiturn import run_test as run_case_07
from .test_cases.case_08_concurrency import run_test as run_case_08
logger = logging.getLogger(__name__)
# 用例注册表
TEST_CASES = {
"case_01_generation": {"func": run_case_01, "name": "纯文本生成基准", "difficulty": "简单", "estimated_seconds": 5},
"case_02_simple_tool": {"func": run_case_02, "name": "简单工具调用延迟", "difficulty": "简单", "estimated_seconds": 3},
"case_03_file_analysis": {"func": run_case_03, "name": "文件读取+分析", "difficulty": "中等", "estimated_seconds": 2},
"case_04_parallel_tool": {"func": run_case_04, "name": "并行工具调用", "difficulty": "中等", "estimated_seconds": 5},
"case_05_long_context": {"func": run_case_05, "name": "长上下文处理", "difficulty": "耗时", "estimated_seconds": 10},
"case_06_reasoning": {"func": run_case_06, "name": "复杂推理任务", "difficulty": "耗时", "estimated_seconds": 15},
"case_07_multiturn": {"func": run_case_07, "name": "多轮对话累积", "difficulty": "耗时", "estimated_seconds": 10},
"case_08_concurrency": {"func": run_case_08, "name": "并发压力测试", "difficulty": "压力", "estimated_seconds": 30},
}
def create_client(model: str):
"""根据配置创建对应的推理客户端"""
if config.LLM_BACKEND == "openai":
from .openai_client import OpenAIClient
return OpenAIClient(model=model)
from .ollama_client import OllamaClient
return OllamaClient(model=model)
# 保留最近 N 次运行结果,超出则删除旧文件
_MAX_RESULTS = 50
def calculate_statistics(results_list: list) -> dict:
"""从用例结果列表计算延迟统计(基于 elapsed_ms)"""
values = [float(r["elapsed_ms"]) for r in results_list if r.get("elapsed_ms") is not None]
if not values:
return {}
pcts = percentiles(values, [50, 95, 99])
return {
"count": len(values),
"mean_ms": round(mean(values), 2),
"std_ms": round(std(values), 2),
"min_ms": round(min_val(values), 2),
"max_ms": round(max_val(values), 2),
"p50_ms": round(pcts[50], 2),
"p95_ms": round(pcts[95], 2),
"p99_ms": round(pcts[99], 2),
}
def calculate_inference_metrics(results_list: list) -> dict:
"""单次遍历聚合推理指标(TTFT / Prefill / Decode / E2E)"""
accumulators = {}
for r in results_list:
for key in ("ttft_ms", "prefill_ms", "e2e_ms"):
val = r.get(key, 0)
if val and float(val) > 0:
accumulators.setdefault(key, []).append(float(val))
decode = r.get("decode_speed_tok_s", 0)
if decode and float(decode) > 0:
accumulators.setdefault("decode_speed_tok_s", []).append(float(decode))
total_tok = r.get("total_tokens", 0)
if total_tok and int(total_tok) > 0:
accumulators.setdefault("total_tokens", []).append(int(total_tok))
result = {}
for key, vals in accumulators.items():
result[f"{key}_mean"] = round(mean(vals), 2)
result[f"{key}_p95"] = round(percentile(vals, 95), 2)
return result
def _cleanup_old_results() -> None:
"""清理过期的运行结果,只保留最近 _MAX_RESULTS 个 run_*.json"""
pattern = os.path.join(config.RESULTS_DIR, "run_*.json")
files = sorted(glob(pattern))
while len(files) > _MAX_RESULTS:
old = files.pop(0)
try:
os.remove(old)
logger.info("清理旧结果: %s", old)
except OSError:
pass
def run_all_tests(cases: list, repeats: int, model: str, skip_heavy: bool,
status: dict = None, stop_event=None) -> dict:
"""运行选定的测试用例,返回完整报告
status: 运行状态字典(由 RunManager 共享,用于向前端汇报进度)
stop_event: threading.Event,置位后停止
"""
client = create_client(model)
# 预检:确认推理服务可用
if not client.ping():
backend = config.LLM_BACKEND
hint = "请确认已运行 'ollama serve' 并拉取模型" if backend == "ollama" else "请检查 OPENAI_BASE_URL / OPENAI_API_KEY 配置"
raise OllamaError(f"无法连接 {backend} 后端服务。{hint}")
# 过滤耗时用例
if skip_heavy:
cases = [c for c in cases if TEST_CASES[c]["difficulty"] != "耗时"]
total_cases = len(cases)
all_results = []
stopped = False
total_start = time.time()
def _update(**kw):
if status is not None:
status.update(kw)
for i, case_id in enumerate(cases, 1):
if stop_event and stop_event.is_set():
stopped = True
_update(message="用户停止测试")
break
case_info = TEST_CASES[case_id]
_update(
status="running",
current_case_index=i - 1,
total_cases=total_cases,
current_case_name=case_info["name"],
progress=round((i - 1) / total_cases * 100, 1),
message=f"正在执行: {case_info['name']} ({case_info['difficulty']})",
)
case_start = time.time()
try:
raw_result = case_info["func"](client=client, repeats=repeats, stop_event=stop_event)
stats = {}
inference_metrics = {}
if raw_result.get("results"):
stats = calculate_statistics(raw_result["results"])
inference_metrics = calculate_inference_metrics(raw_result["results"])
all_results.append({
"case_id": case_id,
"case_name": case_info["name"],
"difficulty": case_info["difficulty"],
"status": "passed",
"raw_data": raw_result,
"statistics": stats,
"inference_metrics": inference_metrics,
"elapsed_seconds": round(time.time() - case_start, 2),
})
_update(
current_case_index=i,
progress=round(i / total_cases * 100, 1),
message=f"完成: {case_info['name']}",
)
logger.info("用例 %s 完成: %s", case_id, case_info["name"])
except Exception as e:
all_results.append({
"case_id": case_id,
"case_name": case_info["name"],
"difficulty": case_info["difficulty"],
"status": "failed",
"error": str(e),
"elapsed_seconds": round(time.time() - case_start, 2),
})
_update(message=f"失败: {case_info['name']} - {e}")
logger.error("用例 %s 失败: %s", case_id, e)
total_elapsed = time.time() - total_start
report = {
"version": "1.0",
"timestamp": datetime.now().isoformat(),
"run_id": datetime.now().strftime("%Y%m%d-%H%M%S"),
"environment": {
"python_version": f"{sys.version_info.major}.{sys.version_info.minor}.{sys.version_info.micro}",
"platform": sys.platform,
"model": model,
"backend": config.LLM_BACKEND,
},
"config": {
"cases_requested": cases,
"repeats": repeats,
"skip_heavy": skip_heavy,
"total_cases": len(cases),
},
"summary": {
"total_cases": len(cases),
"passed": sum(1 for r in all_results if r["status"] == "passed"),
"failed": sum(1 for r in all_results if r["status"] == "failed"),
"stopped": stopped,
"total_elapsed_seconds": round(total_elapsed, 2),
},
"results": all_results,
}
# 保存结果 + 生成 HTML 报告
save_results(report)
return report
def save_results(report: dict) -> None:
"""保存 JSON 结果并生成 HTML 报告"""
os.makedirs(config.RESULTS_DIR, exist_ok=True)
latest_path = os.path.join(config.RESULTS_DIR, "latest.json")
with open(latest_path, "w", encoding="utf-8") as f:
json.dump(report, f, ensure_ascii=False, indent=2)
archive_path = os.path.join(config.RESULTS_DIR, f"run_{report['run_id']}.json")
with open(archive_path, "w", encoding="utf-8") as f:
json.dump(report, f, ensure_ascii=False, indent=2)
# 清理旧结果
_cleanup_old_results()
try:
report_module.generate_report(report)
except Exception as e:
logger.warning("生成 HTML 报告失败: %s", e)