Initial commit of LLM speed test app
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user