Initial commit of LLM speed test app
This commit is contained in:
@@ -0,0 +1,35 @@
|
||||
"""客户端公共基类:推理指标与错误类型
|
||||
|
||||
Ollama 原生客户端与 OpenAI 兼容客户端共享同一数据模型,
|
||||
测试用例通过统一接口(generate/chat)调用,无需关心后端差异。
|
||||
"""
|
||||
from dataclasses import dataclass
|
||||
|
||||
|
||||
class ClientError(Exception):
|
||||
"""推理服务调用相关错误(统一异常基类)"""
|
||||
|
||||
|
||||
@dataclass
|
||||
class InferenceResult:
|
||||
"""单次推理的完整指标(两套后端共用)"""
|
||||
ttft_ms: float = 0.0 # 首 Token 延迟(毫秒)
|
||||
prefill_ms: float = 0.0 # 预填充/处理 prompt 时间(毫秒)
|
||||
decode_speed_tok_s: float = 0.0 # 解码速度(tokens/s)
|
||||
total_tokens: int = 0 # 总 token 数
|
||||
prompt_tokens: int = 0 # prompt token 数
|
||||
completion_tokens: int = 0 # 生成 token 数
|
||||
e2e_ms: float = 0.0 # 端到端耗时(毫秒)
|
||||
response: str = "" # 完整回复文本
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"ttft_ms": round(self.ttft_ms, 2),
|
||||
"prefill_ms": round(self.prefill_ms, 2),
|
||||
"decode_speed_tok_s": round(self.decode_speed_tok_s, 2),
|
||||
"total_tokens": self.total_tokens,
|
||||
"prompt_tokens": self.prompt_tokens,
|
||||
"completion_tokens": self.completion_tokens,
|
||||
"e2e_ms": round(self.e2e_ms, 2),
|
||||
"response_length": len(self.response),
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
"""全局配置"""
|
||||
import os
|
||||
|
||||
# ---- 推理后端选择 ----
|
||||
# 可选值: "ollama"(本地 Ollama)| "openai"(任意 OpenAI 兼容 /v1 服务)
|
||||
LLM_BACKEND = os.environ.get("LLM_BACKEND", "openai")
|
||||
|
||||
# ---- Ollama 服务(backend=ollama)----
|
||||
# 本地 Ollama 默认监听端口 11434,可通过环境变量覆盖
|
||||
OLLAMA_BASE_URL = os.environ.get("OLLAMA_BASE_URL", "http://localhost:11434")
|
||||
# 默认模型(前端会从 /api/models 动态加载实际可用的模型列表)
|
||||
DEFAULT_MODEL = os.environ.get("DEFAULT_MODEL", "qwen3.6:35b-a3b")
|
||||
|
||||
# ---- OpenAI 兼容服务(backend=openai)----
|
||||
OPENAI_BASE_URL = os.environ.get("OPENAI_BASE_URL", "https://ai.lebiztrips.com/v1")
|
||||
OPENAI_API_KEY = os.environ.get("OPENAI_API_KEY", "123")
|
||||
OPENAI_DEFAULT_MODEL = os.environ.get("OPENAI_DEFAULT_MODEL", "Q3.6-35B-A3B-Orig-Thi")
|
||||
|
||||
# ---- 并发压力测试 ----
|
||||
DEFAULT_CONCURRENCY = int(os.environ.get("DEFAULT_CONCURRENCY", "100"))
|
||||
|
||||
# ---- HTTP 服务 ----
|
||||
SERVER_HOST = os.environ.get("SPEED_TEST_HOST", "0.0.0.0")
|
||||
SERVER_PORT = int(os.environ.get("SPEED_TEST_PORT", "8000"))
|
||||
|
||||
# ---- 目录结构 ----
|
||||
ROOT_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
BACKEND_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
FRONTEND_DIR = os.path.join(ROOT_DIR, "frontend")
|
||||
RESULTS_DIR = os.path.join(ROOT_DIR, "results")
|
||||
REPORT_DIR = os.path.join(ROOT_DIR, "report")
|
||||
@@ -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)
|
||||
@@ -0,0 +1,154 @@
|
||||
"""Ollama API 客户端
|
||||
|
||||
通过 HTTP 流式调用本地 Ollama 服务,并测量真实推理指标:
|
||||
- TTFT: 首个 token 到达时间(从发起请求开始计时)
|
||||
- Prefill: prompt 预填充耗时(取自 Ollama 返回的 prompt_eval_duration)
|
||||
- Decode Speed: 解码速度 tokens/s(eval_count / eval_duration)
|
||||
- E2E: 端到端总耗时
|
||||
|
||||
依赖:requests(Ollama 服务需已启动,默认 http://localhost:11434)
|
||||
"""
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
|
||||
import requests
|
||||
|
||||
from . import config
|
||||
from .client_base import InferenceResult, ClientError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class OllamaError(ClientError):
|
||||
"""Ollama 调用相关错误"""
|
||||
|
||||
|
||||
class OllamaClient:
|
||||
"""封装对本地 Ollama 服务的调用"""
|
||||
|
||||
def __init__(self, base_url: str = None, model: str = None, timeout: int = 180):
|
||||
self.base_url = (base_url or config.OLLAMA_BASE_URL).rstrip("/")
|
||||
self.model = model or config.DEFAULT_MODEL
|
||||
self.timeout = timeout
|
||||
self._session = requests.Session()
|
||||
|
||||
# ---------- 服务探测 ----------
|
||||
|
||||
def ping(self) -> bool:
|
||||
"""探测 Ollama 服务是否可用"""
|
||||
try:
|
||||
return self._session.get(f"{self.base_url}/api/tags", timeout=5).ok
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def list_models(self) -> list:
|
||||
"""列出本地已拉取的模型"""
|
||||
try:
|
||||
resp = self._session.get(f"{self.base_url}/api/tags", timeout=5)
|
||||
resp.raise_for_status()
|
||||
return [m["name"] for m in resp.json().get("models", [])]
|
||||
except requests.exceptions.ConnectionError as e:
|
||||
raise OllamaError(
|
||||
f"无法连接 Ollama 服务({self.base_url}),请确认已运行 'ollama serve'"
|
||||
) from e
|
||||
except Exception as e:
|
||||
raise OllamaError(f"获取模型列表失败: {e}") from e
|
||||
|
||||
# ---------- 文本生成(/api/generate)----------
|
||||
|
||||
def generate(self, prompt: str, system: str = None, options: dict = None) -> InferenceResult:
|
||||
"""单轮文本生成"""
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"prompt": prompt,
|
||||
"stream": True,
|
||||
"options": options or {},
|
||||
}
|
||||
if system:
|
||||
payload["system"] = system
|
||||
return self._request(payload)
|
||||
|
||||
# ---------- 多轮对话(/api/chat)----------
|
||||
|
||||
def chat(self, messages: list, options: dict = None) -> InferenceResult:
|
||||
"""多轮对话,messages 为 [{"role": "user"/"assistant", "content": ...}]"""
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"messages": messages,
|
||||
"stream": True,
|
||||
"options": options or {},
|
||||
}
|
||||
return self._request(payload, chat=True)
|
||||
|
||||
# ---------- 核心请求逻辑 ----------
|
||||
|
||||
def _request(self, payload: dict, chat: bool = False) -> InferenceResult:
|
||||
result = InferenceResult()
|
||||
chunks = []
|
||||
final = {}
|
||||
ttft_start = None
|
||||
start = time.perf_counter()
|
||||
|
||||
endpoint = "chat" if chat else "generate"
|
||||
try:
|
||||
resp = self._session.post(
|
||||
f"{self.base_url}/api/{endpoint}",
|
||||
json=payload,
|
||||
stream=True,
|
||||
timeout=self.timeout,
|
||||
)
|
||||
resp.raise_for_status()
|
||||
for line in resp.iter_lines(decode_unicode=True):
|
||||
if not line:
|
||||
continue
|
||||
data = json.loads(line)
|
||||
if chat:
|
||||
piece = (data.get("message") or {}).get("content", "")
|
||||
else:
|
||||
piece = data.get("response", "")
|
||||
done = data.get("done", False)
|
||||
if not done:
|
||||
# 首个非空片段出现时刻记为 TTFT
|
||||
if ttft_start is None and piece:
|
||||
ttft_start = time.perf_counter() - start
|
||||
if piece:
|
||||
chunks.append(piece)
|
||||
else:
|
||||
final = data
|
||||
except requests.exceptions.ConnectionError as e:
|
||||
raise OllamaError(
|
||||
f"无法连接 Ollama 服务({self.base_url}),请确认已运行 'ollama serve' 且模型 {self.model} 已拉取"
|
||||
) from e
|
||||
except requests.exceptions.Timeout as e:
|
||||
raise OllamaError(f"Ollama 请求超时({self.timeout}s): {e}") from e
|
||||
except (json.JSONDecodeError, KeyError) as e:
|
||||
raise OllamaError(f"Ollama 返回数据格式错误: {e}") from e
|
||||
|
||||
result.response = "".join(chunks)
|
||||
result.e2e_ms = (time.perf_counter() - start) * 1000
|
||||
result.ttft_ms = (ttft_start * 1000) if ttft_start else 0.0
|
||||
|
||||
# Ollama 结束块中的指标(单位为纳秒)
|
||||
prompt_eval_count = final.get("prompt_eval_count", 0) or 0
|
||||
prompt_eval_dur_ns = final.get("prompt_eval_duration", 0) or 0
|
||||
eval_count = final.get("eval_count", 0) or 0
|
||||
eval_dur_ns = final.get("eval_duration", 0) or 0
|
||||
|
||||
result.prompt_tokens = int(prompt_eval_count)
|
||||
result.completion_tokens = int(eval_count)
|
||||
result.total_tokens = result.prompt_tokens + result.completion_tokens
|
||||
result.prefill_ms = prompt_eval_dur_ns / 1_000_000
|
||||
if eval_dur_ns:
|
||||
result.decode_speed_tok_s = result.completion_tokens / (eval_dur_ns / 1e9)
|
||||
|
||||
# TTFT 未测到时(如空回复),回退到 prefill 时间
|
||||
if result.ttft_ms <= 0 and result.prefill_ms > 0:
|
||||
result.ttft_ms = result.prefill_ms
|
||||
|
||||
logger.debug(
|
||||
"推理完成: e2e=%.1fms ttft=%.1fms decode=%.1f tok/s tokens=%d/%d",
|
||||
result.e2e_ms, result.ttft_ms, result.decode_speed_tok_s,
|
||||
result.prompt_tokens, result.completion_tokens,
|
||||
)
|
||||
return result
|
||||
@@ -0,0 +1,213 @@
|
||||
"""OpenAI 兼容 API 客户端
|
||||
|
||||
对接任意 OpenAI /v1 兼容服务(OpenAI、Together、硅基流动、llama.cpp server 等),
|
||||
与 OllamaClient 保持相同的 generate/chat 接口,测试用例无需区分后端。
|
||||
|
||||
指标来源:
|
||||
- TTFT: 流式请求中首个 content 片段到达时间(真实测量)
|
||||
- Prefill: 优先取服务端 timings.prompt_ms(llama.cpp server 暴露),否则回退 TTFT
|
||||
- Decode Speed: 优先取服务端 timings.predicted_per_second,否则按 completion_tokens/e2e 估算
|
||||
- E2E: 客户端请求总耗时
|
||||
"""
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
|
||||
import requests
|
||||
|
||||
from . import config
|
||||
from .client_base import InferenceResult, ClientError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class OpenAIError(ClientError):
|
||||
"""OpenAI 兼容 API 调用错误"""
|
||||
|
||||
|
||||
class OpenAIClient:
|
||||
"""封装对 OpenAI 兼容 /v1 服务的调用"""
|
||||
|
||||
def __init__(self, base_url: str = None, model: str = None, api_key: str = None,
|
||||
timeout: int = 300):
|
||||
self.base_url = (base_url or config.OPENAI_BASE_URL).rstrip("/")
|
||||
self.model = model or config.OPENAI_DEFAULT_MODEL
|
||||
self.api_key = api_key if api_key not in (None, "") else config.OPENAI_API_KEY
|
||||
self.timeout = timeout
|
||||
self._session = requests.Session()
|
||||
|
||||
# ---------- 请求头 ----------
|
||||
|
||||
@property
|
||||
def _headers(self) -> dict:
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if self.api_key:
|
||||
headers["Authorization"] = f"Bearer {self.api_key}"
|
||||
return headers
|
||||
|
||||
# ---------- 服务探测 ----------
|
||||
|
||||
def ping(self) -> bool:
|
||||
"""探测服务是否可用"""
|
||||
try:
|
||||
return self._session.get(
|
||||
f"{self.base_url}/models", headers=self._headers, timeout=5
|
||||
).ok
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def list_models(self) -> list:
|
||||
"""列出可用模型"""
|
||||
try:
|
||||
resp = self._session.get(
|
||||
f"{self.base_url}/models", headers=self._headers, timeout=10
|
||||
)
|
||||
resp.raise_for_status()
|
||||
data = resp.json().get("data", [])
|
||||
models = []
|
||||
for m in data:
|
||||
mid = m.get("id") or m.get("name")
|
||||
if mid:
|
||||
models.append(mid)
|
||||
return models
|
||||
except requests.exceptions.ConnectionError as e:
|
||||
raise OpenAIError(
|
||||
f"无法连接服务({self.base_url}),请检查地址与网络"
|
||||
) from e
|
||||
except requests.exceptions.HTTPError as e:
|
||||
if resp.status_code in (401, 403):
|
||||
raise OpenAIError("API Key 无效或被拒绝(401/403)") from e
|
||||
raise OpenAIError(f"获取模型列表失败: {resp.status_code} {resp.text[:200]}") from e
|
||||
except Exception as e:
|
||||
raise OpenAIError(f"获取模型列表失败: {e}") from e
|
||||
|
||||
# ---------- 文本生成 ----------
|
||||
|
||||
def generate(self, prompt: str, system: str = None, options: dict = None) -> InferenceResult:
|
||||
"""单轮文本生成(流式),等价于单条 user 消息的 chat"""
|
||||
messages = []
|
||||
if system:
|
||||
messages.append({"role": "system", "content": system})
|
||||
messages.append({"role": "user", "content": prompt})
|
||||
return self.chat(messages, options)
|
||||
|
||||
# ---------- 多轮对话 ----------
|
||||
|
||||
def chat(self, messages: list, options: dict = None) -> InferenceResult:
|
||||
"""多轮对话,messages 为 [{"role": ..., "content": ...}]"""
|
||||
options = options or {}
|
||||
payload = {
|
||||
"model": self.model,
|
||||
"messages": messages,
|
||||
"stream": True,
|
||||
"temperature": options.get("temperature", 0.3),
|
||||
"max_tokens": options.get("num_predict", 256),
|
||||
}
|
||||
|
||||
start = time.perf_counter()
|
||||
ttft_ms = 0.0
|
||||
chunks = []
|
||||
final = {}
|
||||
timings = {}
|
||||
|
||||
try:
|
||||
resp = self._session.post(
|
||||
f"{self.base_url}/chat/completions",
|
||||
headers=self._headers,
|
||||
json=payload,
|
||||
stream=True,
|
||||
timeout=self.timeout,
|
||||
)
|
||||
except requests.exceptions.ConnectionError as e:
|
||||
raise OpenAIError(
|
||||
f"无法连接服务({self.base_url}),请检查地址与网络"
|
||||
) from e
|
||||
except requests.exceptions.Timeout as e:
|
||||
raise OpenAIError(f"请求超时({self.timeout}s): {e}") from e
|
||||
|
||||
if resp.status_code != 200:
|
||||
err_text = resp.text[:300]
|
||||
raise OpenAIError(f"API 错误 {resp.status_code}: {err_text}")
|
||||
|
||||
# 流式解析 SSE
|
||||
try:
|
||||
for line in resp.iter_lines(decode_unicode=True):
|
||||
if not line:
|
||||
continue
|
||||
line = line.strip()
|
||||
if line.startswith("data:"):
|
||||
line = line[5:].strip()
|
||||
if line == "[DONE]":
|
||||
break
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
data = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
|
||||
choices = data.get("choices") or []
|
||||
if not choices:
|
||||
continue
|
||||
delta = choices[0].get("delta") or {}
|
||||
# 兼容推理模型:token 可能放在 reasoning_content(思考过程)或 content
|
||||
piece = delta.get("content") or delta.get("reasoning_content") or ""
|
||||
|
||||
if piece:
|
||||
# 首个内容片段到达时间 = TTFT
|
||||
if ttft_ms == 0.0:
|
||||
ttft_ms = (time.perf_counter() - start) * 1000
|
||||
chunks.append(piece)
|
||||
|
||||
# 服务端最终块可能携带 usage / timings
|
||||
if choices[0].get("finish_reason"):
|
||||
final = data
|
||||
if data.get("usage"):
|
||||
final = data
|
||||
except (requests.exceptions.ConnectionError, requests.exceptions.ChunkedEncodingError) as e:
|
||||
raise OpenAIError(f"流式读取中断: {e}") from e
|
||||
|
||||
e2e_ms = (time.perf_counter() - start) * 1000
|
||||
result = InferenceResult()
|
||||
result.response = "".join(chunks)
|
||||
result.e2e_ms = e2e_ms
|
||||
result.ttft_ms = ttft_ms
|
||||
|
||||
# usage: token 计数(部分服务端流式响应不含 usage,回退用 timings)
|
||||
usage = final.get("usage", {}) or {}
|
||||
prompt_tok = int(usage.get("prompt_tokens", 0) or 0)
|
||||
completion_tok = int(usage.get("completion_tokens", 0) or 0)
|
||||
|
||||
# timings: llama.cpp / llama-server 暴露的服务端计时(单位毫秒)
|
||||
timings = final.get("timings", {}) or {}
|
||||
prompt_n = int(timings.get("prompt_n", 0) or 0)
|
||||
predicted_n = int(timings.get("predicted_n", 0) or 0)
|
||||
prompt_ms = float(timings.get("prompt_ms", 0) or 0)
|
||||
predicted_ms = float(timings.get("predicted_ms", 0) or 0)
|
||||
predicted_pps = float(timings.get("predicted_per_second", 0) or 0)
|
||||
|
||||
result.prompt_tokens = prompt_tok or prompt_n
|
||||
result.completion_tokens = completion_tok or predicted_n
|
||||
result.total_tokens = result.prompt_tokens + result.completion_tokens
|
||||
|
||||
result.prefill_ms = prompt_ms
|
||||
|
||||
# 解码速度:优先服务端计时
|
||||
if predicted_pps > 0:
|
||||
result.decode_speed_tok_s = predicted_pps
|
||||
elif result.completion_tokens > 0 and result.e2e_ms > 0:
|
||||
result.decode_speed_tok_s = result.completion_tokens / (result.e2e_ms / 1000)
|
||||
|
||||
# TTFT 回退:流式被缓冲或未测到时,用服务端 prompt_ms
|
||||
if result.ttft_ms <= 0 or result.ttft_ms > e2e_ms * 0.95:
|
||||
if prompt_ms > 0:
|
||||
result.ttft_ms = prompt_ms
|
||||
else:
|
||||
result.ttft_ms = e2e_ms if e2e_ms > 0 else 0.0
|
||||
|
||||
logger.debug(
|
||||
"推理完成(openai): e2e=%.1fms ttft=%.1fms decode=%.1f tok/s tokens=%d/%d",
|
||||
result.e2e_ms, result.ttft_ms, result.decode_speed_tok_s,
|
||||
result.prompt_tokens, result.completion_tokens,
|
||||
)
|
||||
return result
|
||||
@@ -0,0 +1,279 @@
|
||||
"""生成 HTML 可视化报告"""
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
|
||||
from . import config
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def generate_report(report: dict) -> str:
|
||||
"""从报告数据生成 HTML 报告文件,返回输出路径"""
|
||||
html = build_html(report)
|
||||
|
||||
os.makedirs(config.REPORT_DIR, exist_ok=True)
|
||||
output_path = os.path.join(config.REPORT_DIR, "report.html")
|
||||
with open(output_path, "w", encoding="utf-8") as f:
|
||||
f.write(html)
|
||||
|
||||
logger.info("报告已生成: %s", output_path)
|
||||
return output_path
|
||||
|
||||
|
||||
def _diff_color(difficulty: str) -> str:
|
||||
"""难度对应的颜色"""
|
||||
return {
|
||||
"简单": "#22c55e",
|
||||
"中等": "#f59e0b",
|
||||
"耗时": "#ef4444",
|
||||
"压力": "#8b5cf6",
|
||||
}.get(difficulty, "#6b7280")
|
||||
|
||||
|
||||
def _stats_rows(results: list) -> str:
|
||||
"""构建详细数据表格行 HTML"""
|
||||
rows = []
|
||||
for r in results:
|
||||
stats = r.get("statistics", {})
|
||||
if not stats:
|
||||
continue
|
||||
rows.append(
|
||||
f'<tr>'
|
||||
f'<td><strong>{r["case_name"]}</strong></td>'
|
||||
f'<td><span class="diff-badge" style="background:{_diff_color(r["difficulty"])}">{r["difficulty"]}</span></td>'
|
||||
f'<td class="status-{r["status"]}" style="font-weight:bold">{r["status"]}</td>'
|
||||
f'<td class="mono">{stats.get("mean_ms", "N/A")}</td>'
|
||||
f'<td class="mono">{stats.get("std_ms", "N/A")}</td>'
|
||||
f'<td class="mono">{stats.get("min_ms", "N/A")}</td>'
|
||||
f'<td class="mono">{stats.get("max_ms", "N/A")}</td>'
|
||||
f'<td class="mono">{stats.get("p95_ms", "N/A")}</td>'
|
||||
f'</tr>'
|
||||
)
|
||||
return "\n".join(rows)
|
||||
|
||||
|
||||
def _chart_data(results: list) -> tuple:
|
||||
"""提取图表数据:(labels, means, p95s)"""
|
||||
labels, means, p95s = [], [], []
|
||||
for r in results:
|
||||
stats = r.get("statistics", {})
|
||||
if stats:
|
||||
labels.append(r["case_name"])
|
||||
means.append(stats.get("mean_ms", 0))
|
||||
p95s.append(stats.get("p95_ms", 0))
|
||||
return labels, means, p95s
|
||||
|
||||
|
||||
def _performance_analysis(results: list) -> str:
|
||||
"""性能分析摘要"""
|
||||
_, means, _ = _chart_data(results)
|
||||
if len(means) < 2:
|
||||
return ""
|
||||
|
||||
best_idx = means.index(min(means))
|
||||
worst_idx = means.index(max(means))
|
||||
ratio = (means[worst_idx] / means[best_idx]) if means[best_idx] > 0 else "N/A"
|
||||
avg = sum(means) / len(means)
|
||||
|
||||
return (
|
||||
f'<div class="analysis-box">'
|
||||
f'<h3>性能分析</h3>'
|
||||
f'<ul>'
|
||||
f'<li><strong>最快用例:</strong> {results[best_idx]["case_name"]} ({means[best_idx]:.2f}ms)</li>'
|
||||
f'<li><strong>最慢用例:</strong> {results[worst_idx]["case_name"]} ({means[worst_idx]:.2f}ms)</li>'
|
||||
f'<li><strong>最快/最慢比:</strong> {ratio}x</li>'
|
||||
f'<li><strong>平均延迟:</strong> {avg:.2f}ms</li>'
|
||||
f'</ul></div>'
|
||||
)
|
||||
|
||||
|
||||
def _concurrency_summary(results: list) -> str:
|
||||
"""提取并发压力测试摘要(case_08),无则返回空串"""
|
||||
for r in results:
|
||||
raw = r.get("raw_data", {}) or {}
|
||||
if r.get("case_id") == "case_08_concurrency" and raw:
|
||||
def g(key, suffix=""):
|
||||
val = raw.get(key)
|
||||
return "N/A" if val is None else f"{val}{suffix}"
|
||||
|
||||
keep = raw.get("tps_keep_rate_pct", 0)
|
||||
keep_color = "#22c55e" if keep >= 60 else ("#f59e0b" if keep > 0 else "#ef4444")
|
||||
err = raw.get("error_rate_pct", 0)
|
||||
err_color = "#22c55e" if err == 0 else "#ef4444"
|
||||
|
||||
return (
|
||||
f'<div class="section">'
|
||||
f'<h2>⚡ 并发压力测试摘要</h2>'
|
||||
f'<div class="dashboard">'
|
||||
f'<div class="stat-card"><div class="label">并发数</div><div class="value" style="color:#8b5cf6">{g("concurrency")}</div></div>'
|
||||
f'<div class="stat-card"><div class="label">基准 TPS</div><div class="value" style="color:#6366f1">{g("baseline_tps")}</div></div>'
|
||||
f'<div class="stat-card"><div class="label">并发 TPS</div><div class="value" style="color:#6366f1">{g("concurrent_tps")}</div></div>'
|
||||
f'<div class="stat-card"><div class="label">TPS 保持率</div><div class="value" style="color:{keep_color}">{g("tps_keep_rate_pct", "%")}</div></div>'
|
||||
f'<div class="stat-card"><div class="label">P95 延迟</div><div class="value" style="color:#ef4444">{g("e2e_p95_ms", "ms")}</div></div>'
|
||||
f'<div class="stat-card"><div class="label">P99 延迟</div><div class="value" style="color:#ef4444">{g("e2e_p99_ms", "ms")}</div></div>'
|
||||
f'<div class="stat-card"><div class="label">错误率</div><div class="value" style="color:{err_color}">{g("error_rate_pct", "%")}</div></div>'
|
||||
f'</div></div>'
|
||||
)
|
||||
return ""
|
||||
|
||||
|
||||
def build_html(report: dict) -> str:
|
||||
"""构建 HTML 报告"""
|
||||
summary = report.get("summary", {})
|
||||
results = report.get("results", [])
|
||||
env = report.get("environment", {})
|
||||
config_ = report.get("config", {})
|
||||
timestamp = report.get("timestamp", "")
|
||||
run_id = report.get("run_id", "")
|
||||
|
||||
case_names, case_means, case_p95s = _chart_data(results)
|
||||
|
||||
return f"""<!DOCTYPE html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||
<title>大模型速度测试报告 - {run_id}</title>
|
||||
<script src="https://cdn.jsdelivr.net/npm/chart.js@4.4.0/dist/chart.umd.min.js"></script>
|
||||
<style>
|
||||
* {{ margin: 0; padding: 0; box-sizing: border-box; }}
|
||||
body {{
|
||||
font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', 'PingFang SC', 'Microsoft YaHei', sans-serif;
|
||||
background: #f8fafc; color: #1e293b; line-height: 1.6;
|
||||
}}
|
||||
.container {{ max-width: 1400px; margin: 0 auto; padding: 20px; }}
|
||||
.header {{ background: linear-gradient(135deg, #667eea 0%, #764ba2 100%); color: white; padding: 30px; border-radius: 12px; margin-bottom: 20px; }}
|
||||
.header h1 {{ font-size: 28px; margin-bottom: 10px; }}
|
||||
.header .meta {{ font-size: 14px; opacity: 0.9; }}
|
||||
.dashboard {{ display: grid; grid-template-columns: repeat(auto-fit, minmax(200px, 1fr)); gap: 15px; margin-bottom: 20px; }}
|
||||
.stat-card {{ background: white; padding: 20px; border-radius: 10px; box-shadow: 0 2px 8px rgba(0,0,0,0.1); }}
|
||||
.stat-card .label {{ font-size: 12px; color: #64748b; text-transform: uppercase; }}
|
||||
.stat-card .value {{ font-size: 28px; font-weight: bold; margin-top: 5px; }}
|
||||
.section {{ background: white; border-radius: 10px; padding: 25px; margin-bottom: 20px; box-shadow: 0 2px 8px rgba(0,0,0,0.1); }}
|
||||
.section h2 {{ font-size: 20px; margin-bottom: 15px; padding-bottom: 10px; border-bottom: 2px solid #e2e8f0; }}
|
||||
table {{ width: 100%; border-collapse: collapse; }}
|
||||
th, td {{ padding: 12px; text-align: left; border-bottom: 1px solid #e2e8f0; }}
|
||||
th {{ background: #f1f5f9; font-weight: 600; font-size: 13px; text-transform: uppercase; }}
|
||||
tr:hover {{ background: #f8fafc; }}
|
||||
.mono {{ font-family: 'SF Mono', 'Fira Code', monospace; }}
|
||||
.diff-badge {{ padding: 4px 10px; border-radius: 12px; color: white; font-size: 12px; font-weight: 600; }}
|
||||
.status-passed {{ color: #22c55e; }}
|
||||
.status-failed {{ color: #ef4444; }}
|
||||
.chart-container {{ position: relative; height: 400px; margin: 20px 0; }}
|
||||
.analysis-box {{ background: #f0fdf4; border-left: 4px solid #22c55e; padding: 15px; border-radius: 6px; }}
|
||||
.analysis-box ul {{ padding-left: 20px; margin-top: 10px; }}
|
||||
.analysis-box li {{ margin: 5px 0; }}
|
||||
footer {{ text-align: center; padding: 20px; color: #64748b; font-size: 13px; }}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div class="container">
|
||||
<div class="header">
|
||||
<h1>大模型速度测试报告</h1>
|
||||
<div class="meta">
|
||||
Run ID: {run_id} | 时间: {timestamp} | 用例: {config_.get('total_cases', 0)} | 重复: {config_.get('repeats', 0)}次
|
||||
<br>模型: {env.get('model', 'N/A')} | 平台: {env.get('platform', 'N/A')}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="dashboard">
|
||||
<div class="stat-card"><div class="label">总用例数</div><div class="value">{summary.get('total_cases', 0)}</div></div>
|
||||
<div class="stat-card"><div class="label">通过</div><div class="value" style="color:#22c55e">{summary.get('passed', 0)}</div></div>
|
||||
<div class="stat-card"><div class="label">失败</div><div class="value" style="color:#ef4444">{summary.get('failed', 0)}</div></div>
|
||||
<div class="stat-card"><div class="label">总耗时</div><div class="value">{summary.get('total_elapsed_seconds', 0):.2f}s</div></div>
|
||||
</div>
|
||||
|
||||
{_concurrency_summary(results)}
|
||||
|
||||
{_performance_analysis(results)}
|
||||
|
||||
<div class="section">
|
||||
<h2>性能对比柱状图</h2>
|
||||
<div class="chart-container"><canvas id="barChart"></canvas></div>
|
||||
</div>
|
||||
|
||||
<div class="section">
|
||||
<h2>P95延迟对比</h2>
|
||||
<div class="chart-container"><canvas id="p95Chart"></canvas></div>
|
||||
</div>
|
||||
|
||||
<div class="section">
|
||||
<h2>详细数据</h2>
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>用例名称</th><th>难度</th><th>状态</th><th>平均延迟 (ms)</th>
|
||||
<th>标准差</th><th>最小</th><th>最大</th><th>P95</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
{_stats_rows(results)}
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
|
||||
<div class="section">
|
||||
<h2>环境信息</h2>
|
||||
<table>
|
||||
<tr><td>模型</td><td class="mono">{env.get('model', 'N/A')}</td></tr>
|
||||
<tr><td>运行平台</td><td class="mono">{env.get('platform', 'N/A')}</td></tr>
|
||||
<tr><td>运行ID</td><td class="mono">{run_id}</td></tr>
|
||||
<tr><td>测试时间</td><td class="mono">{timestamp}</td></tr>
|
||||
</table>
|
||||
</div>
|
||||
|
||||
<footer>大模型速度测试助手 | 报告由 report.py 自动生成</footer>
|
||||
</div>
|
||||
|
||||
<script>
|
||||
const barCtx = document.getElementById('barChart').getContext('2d');
|
||||
new Chart(barCtx, {{
|
||||
type: 'bar',
|
||||
data: {{
|
||||
labels: {json.dumps(case_names, ensure_ascii=False)},
|
||||
datasets: [{{
|
||||
label: '平均延迟 (ms)',
|
||||
data: {json.dumps(case_means)},
|
||||
backgroundColor: 'rgba(102, 126, 234, 0.7)',
|
||||
borderColor: 'rgba(102, 126, 234, 1)',
|
||||
borderWidth: 1
|
||||
}}]
|
||||
}},
|
||||
options: {{
|
||||
responsive: true,
|
||||
maintainAspectRatio: false,
|
||||
plugins: {{
|
||||
legend: {{ display: false }},
|
||||
title: {{ display: true, text: '各用例平均延迟对比 (ms)', font: {{ size: 16 }} }}
|
||||
}},
|
||||
scales: {{ y: {{ beginAtZero: true, title: {{ display: true, text: '延迟 (ms)' }} }} }}
|
||||
}}
|
||||
}});
|
||||
|
||||
const p95Ctx = document.getElementById('p95Chart').getContext('2d');
|
||||
new Chart(p95Ctx, {{
|
||||
type: 'bar',
|
||||
data: {{
|
||||
labels: {json.dumps(case_names, ensure_ascii=False)},
|
||||
datasets: [{{
|
||||
label: 'P95延迟 (ms)',
|
||||
data: {json.dumps(case_p95s)},
|
||||
backgroundColor: 'rgba(239, 68, 68, 0.7)',
|
||||
borderColor: 'rgba(239, 68, 68, 1)',
|
||||
borderWidth: 1
|
||||
}}]
|
||||
}},
|
||||
options: {{
|
||||
responsive: true,
|
||||
maintainAspectRatio: false,
|
||||
plugins: {{
|
||||
legend: {{ display: false }},
|
||||
title: {{ display: true, text: '各用例P95延迟 (ms)', font: {{ size: 16 }} }}
|
||||
}},
|
||||
scales: {{ y: {{ beginAtZero: true, title: {{ display: true, text: '延迟 (ms)' }} }} }}
|
||||
}}
|
||||
}});
|
||||
</script>
|
||||
</body>
|
||||
</html>"""
|
||||
@@ -0,0 +1,380 @@
|
||||
"""HTTP 服务器:提供 API + 托管前端静态文件
|
||||
|
||||
启动后访问 http://localhost:8000 即可打开测试界面。
|
||||
"""
|
||||
import json
|
||||
import logging
|
||||
import mimetypes
|
||||
import os
|
||||
import threading
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from datetime import datetime
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from urllib.parse import urlparse
|
||||
|
||||
from . import config
|
||||
from .client_base import ClientError
|
||||
from .engine import TEST_CASES, run_all_tests
|
||||
from .ollama_client import OllamaClient, OllamaError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class RunState:
|
||||
"""测试运行状态(线程安全:所有读写通过 RunManager 加锁)"""
|
||||
status: str = "idle" # idle | running | done | stopped | failed
|
||||
run_id: str = ""
|
||||
model: str = ""
|
||||
repeats: int = 0
|
||||
skip_heavy: bool = False
|
||||
cases: list = field(default_factory=list)
|
||||
current_case_index: int = 0
|
||||
total_cases: int = 0
|
||||
current_case_name: str = ""
|
||||
message: str = "准备中"
|
||||
progress: float = 0.0
|
||||
started_at: str = ""
|
||||
finished_at: str = ""
|
||||
summary: dict = field(default_factory=dict)
|
||||
|
||||
def update(self, values=None, **kwargs) -> None:
|
||||
"""批量更新字段(引擎通过 status.update(kw) 或 status.update(**kw) 汇报进度)"""
|
||||
updates = dict(values or {})
|
||||
updates.update(kwargs)
|
||||
for key, value in updates.items():
|
||||
if hasattr(self, key):
|
||||
setattr(self, key, value)
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
class RunManager:
|
||||
"""管理测试运行状态(单实例、同一时间只允许一次运行)"""
|
||||
|
||||
def __init__(self):
|
||||
self._lock = threading.Lock()
|
||||
self._stop_event = threading.Event()
|
||||
self._state = RunState()
|
||||
self._thread = None
|
||||
|
||||
@property
|
||||
def is_running(self) -> bool:
|
||||
return self._state.status == "running"
|
||||
|
||||
def start(self, cases: list, repeats: int, model: str, skip_heavy: bool) -> str:
|
||||
with self._lock:
|
||||
if self.is_running:
|
||||
raise RuntimeError("已有测试正在运行,请等待完成或先停止")
|
||||
self._stop_event.clear()
|
||||
run_id = datetime.now().strftime("%Y%m%d-%H%M%S")
|
||||
self._state = RunState(
|
||||
status="running",
|
||||
run_id=run_id,
|
||||
model=model,
|
||||
repeats=repeats,
|
||||
skip_heavy=skip_heavy,
|
||||
cases=list(cases),
|
||||
total_cases=len(cases),
|
||||
started_at=datetime.now().isoformat(),
|
||||
)
|
||||
self._thread = threading.Thread(
|
||||
target=self._run, args=(cases, repeats, model, skip_heavy), daemon=True
|
||||
)
|
||||
self._thread.start()
|
||||
return run_id
|
||||
|
||||
def stop(self) -> None:
|
||||
"""请求停止(置位事件,引擎在用例之间检查)"""
|
||||
self._stop_event.set()
|
||||
|
||||
def status(self) -> dict:
|
||||
"""返回当前运行状态(对外只读拷贝)"""
|
||||
with self._lock:
|
||||
return self._state.to_dict()
|
||||
|
||||
def _run(self, cases, repeats, model, skip_heavy):
|
||||
try:
|
||||
report = run_all_tests(
|
||||
cases=cases,
|
||||
repeats=repeats,
|
||||
model=model,
|
||||
skip_heavy=skip_heavy,
|
||||
status=self._state,
|
||||
stop_event=self._stop_event,
|
||||
)
|
||||
stopped = report["summary"].get("stopped", False)
|
||||
with self._lock:
|
||||
self._state.status = "stopped" if stopped else "done"
|
||||
self._state.progress = 100
|
||||
self._state.message = "测试已停止" if stopped else "测试完成"
|
||||
self._state.finished_at = datetime.now().isoformat()
|
||||
self._state.summary = report["summary"]
|
||||
except Exception as e:
|
||||
logger.exception("测试引擎异常")
|
||||
with self._lock:
|
||||
self._state.status = "failed"
|
||||
self._state.message = f"测试失败: {e}"
|
||||
self._state.finished_at = datetime.now().isoformat()
|
||||
|
||||
|
||||
manager = RunManager()
|
||||
|
||||
|
||||
class Handler(BaseHTTPRequestHandler):
|
||||
"""API + 静态文件处理"""
|
||||
|
||||
# ---------- HTTP 基础 ----------
|
||||
|
||||
def _send_json(self, obj, code=200):
|
||||
body = json.dumps(obj, ensure_ascii=False).encode("utf-8")
|
||||
self.send_response(code)
|
||||
self.send_header("Content-Type", "application/json; charset=utf-8")
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
|
||||
def _read_json(self) -> dict:
|
||||
try:
|
||||
length = int(self.headers.get("Content-Length", 0) or 0)
|
||||
except ValueError:
|
||||
length = 0
|
||||
if length <= 0:
|
||||
return {}
|
||||
try:
|
||||
return json.loads(self.rfile.read(length).decode("utf-8"))
|
||||
except json.JSONDecodeError:
|
||||
return {}
|
||||
|
||||
def log_message(self, format, *args):
|
||||
print(f"[API] {args[0]}")
|
||||
|
||||
# ---------- GET ----------
|
||||
|
||||
def do_GET(self):
|
||||
path = urlparse(self.path).path
|
||||
if path == "/api/health":
|
||||
self._send_json({"status": "ok"})
|
||||
elif path == "/api/config":
|
||||
self._handle_config()
|
||||
elif path == "/api/models":
|
||||
self._handle_models()
|
||||
elif path == "/api/cases":
|
||||
cases = {
|
||||
cid: {
|
||||
"name": info["name"],
|
||||
"difficulty": info["difficulty"],
|
||||
"estimated_seconds": info["estimated_seconds"],
|
||||
}
|
||||
for cid, info in TEST_CASES.items()
|
||||
}
|
||||
self._send_json(cases)
|
||||
elif path == "/api/status":
|
||||
self._send_json(manager.status())
|
||||
elif path == "/api/results/latest.json":
|
||||
self._handle_latest()
|
||||
elif path == "/api/results/list":
|
||||
self._handle_results_list()
|
||||
elif path.startswith("/api/results/"):
|
||||
self._handle_result_file(path)
|
||||
elif path == "/report.html":
|
||||
self._serve_file(os.path.join(config.REPORT_DIR, "report.html"))
|
||||
else:
|
||||
self._serve_static(path)
|
||||
|
||||
# ---------- POST ----------
|
||||
|
||||
def do_POST(self):
|
||||
path = urlparse(self.path).path
|
||||
body = self._read_json()
|
||||
if path == "/api/run":
|
||||
self._handle_run(body)
|
||||
elif path == "/api/stop":
|
||||
manager.stop()
|
||||
self._send_json({"ok": True, "message": "已请求停止"})
|
||||
else:
|
||||
self._send_json({"error": "Not found"}, 404)
|
||||
|
||||
# ---------- API 处理器 ----------
|
||||
|
||||
def _handle_config(self):
|
||||
self._send_json({
|
||||
"backend": config.LLM_BACKEND,
|
||||
"ollama_base_url": config.OLLAMA_BASE_URL,
|
||||
"openai_base_url": config.OPENAI_BASE_URL,
|
||||
"default_model": config.DEFAULT_MODEL,
|
||||
"openai_default_model": config.OPENAI_DEFAULT_MODEL,
|
||||
"concurrency": config.DEFAULT_CONCURRENCY,
|
||||
})
|
||||
|
||||
def _handle_models(self):
|
||||
try:
|
||||
if config.LLM_BACKEND == "openai":
|
||||
from .openai_client import OpenAIClient
|
||||
models = OpenAIClient().list_models()
|
||||
else:
|
||||
models = OllamaClient().list_models()
|
||||
self._send_json({"models": models, "error": None})
|
||||
except ClientError as e:
|
||||
self._send_json({"models": [], "error": str(e)})
|
||||
|
||||
def _handle_run(self, body):
|
||||
cases = body.get("cases") or []
|
||||
repeats = int(body.get("repeats", 3) or 3)
|
||||
model = body.get("model") or config.DEFAULT_MODEL
|
||||
skip_heavy = bool(body.get("skip_heavy", False))
|
||||
|
||||
# 校验用例
|
||||
unknown = [c for c in cases if c not in TEST_CASES]
|
||||
if not cases:
|
||||
self._send_json({"error": "未选择测试用例"}, 400)
|
||||
return
|
||||
if unknown:
|
||||
self._send_json({"error": f"未知用例: {', '.join(unknown)}"}, 400)
|
||||
return
|
||||
if repeats < 1 or repeats > 50:
|
||||
self._send_json({"error": "重复次数需在 1-50 之间"}, 400)
|
||||
return
|
||||
|
||||
try:
|
||||
run_id = manager.start(cases, repeats, model, skip_heavy)
|
||||
self._send_json({"run_id": run_id})
|
||||
except RuntimeError as e:
|
||||
self._send_json({"error": str(e)}, 409)
|
||||
|
||||
def _handle_latest(self):
|
||||
"""返回最新测试结果,文件损坏时优雅降级"""
|
||||
latest_path = os.path.join(config.RESULTS_DIR, "latest.json")
|
||||
if not os.path.exists(latest_path):
|
||||
self._send_json({"error": "暂无测试结果,请先运行测试"}, 404)
|
||||
return
|
||||
try:
|
||||
with open(latest_path, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
self._send_json(data)
|
||||
except (json.JSONDecodeError, OSError) as e:
|
||||
logger.error("读取 latest.json 失败: %s", e)
|
||||
self._send_json({"error": "结果文件损坏,请先运行新测试"}, 500)
|
||||
|
||||
def _handle_results_list(self):
|
||||
"""列出所有 run_*.json 结果文件"""
|
||||
import glob as glob_mod
|
||||
pattern = os.path.join(config.RESULTS_DIR, "run_*.json")
|
||||
files = sorted(glob_mod.glob(pattern))
|
||||
result = []
|
||||
for f in files:
|
||||
name = os.path.basename(f)
|
||||
try:
|
||||
with open(f, "r", encoding="utf-8") as fh:
|
||||
data = json.load(fh)
|
||||
summary = data.get("summary", {})
|
||||
result.append({
|
||||
"filename": name,
|
||||
"timestamp": data.get("timestamp", ""),
|
||||
"model": data.get("environment", {}).get("model", ""),
|
||||
"total_cases": summary.get("total_cases", 0),
|
||||
"passed": summary.get("passed", 0),
|
||||
"failed": summary.get("failed", 0),
|
||||
"elapsed_seconds": summary.get("total_elapsed_seconds", 0),
|
||||
})
|
||||
except (json.JSONDecodeError, OSError):
|
||||
result.append({"filename": name, "error": "无法读取"})
|
||||
self._send_json({"results": result})
|
||||
|
||||
def _handle_result_file(self, path):
|
||||
"""返回单个历史结果文件 /api/results/<filename>"""
|
||||
filename = path.rsplit("/", 1)[-1]
|
||||
if not filename.endswith(".json") or "/" in filename or "\\" in filename:
|
||||
self._send_json({"error": "Forbidden"}, 403)
|
||||
return
|
||||
full = os.path.join(config.RESULTS_DIR, filename)
|
||||
if not os.path.isfile(full):
|
||||
self._send_json({"error": "Not found"}, 404)
|
||||
return
|
||||
try:
|
||||
with open(full, "r", encoding="utf-8") as f:
|
||||
data = json.load(f)
|
||||
self._send_json(data)
|
||||
except (json.JSONDecodeError, OSError):
|
||||
self._send_json({"error": "文件损坏"}, 500)
|
||||
|
||||
# ---------- 静态文件 ----------
|
||||
|
||||
def _serve_static(self, path):
|
||||
# 只允许从 frontend/ 目录读取
|
||||
if path == "/" or path == "/index.html":
|
||||
rel = "index.html"
|
||||
else:
|
||||
rel = path.lstrip("/")
|
||||
|
||||
full = os.path.normpath(os.path.join(config.FRONTEND_DIR, rel))
|
||||
frontend_root = os.path.normpath(config.FRONTEND_DIR)
|
||||
|
||||
# 防止路径遍历攻击
|
||||
if not os.path.abspath(full).startswith(os.path.abspath(frontend_root + os.sep)):
|
||||
self._send_json({"error": "Forbidden"}, 403)
|
||||
return
|
||||
|
||||
# 防止符号链接指向外部目录
|
||||
# 逐段检查路径中每个组件是否安全
|
||||
parts = os.path.relpath(full, frontend_root).split(os.sep)
|
||||
current = frontend_root
|
||||
for part in parts:
|
||||
current = os.path.join(current, part)
|
||||
if os.path.islink(current):
|
||||
link_target = os.path.realpath(current)
|
||||
if not link_target.startswith(os.path.realpath(frontend_root) + os.sep):
|
||||
self._send_json({"error": "Forbidden"}, 403)
|
||||
return
|
||||
if not os.path.exists(current):
|
||||
break
|
||||
|
||||
self._serve_file(full)
|
||||
|
||||
def _serve_file(self, full_path):
|
||||
if not os.path.isfile(full_path):
|
||||
self._send_json({"error": "Not found"}, 404)
|
||||
return
|
||||
content_type = mimetypes.guess_type(full_path)[0] or "application/octet-stream"
|
||||
with open(full_path, "rb") as f:
|
||||
data = f.read()
|
||||
self.send_response(200)
|
||||
self.send_header("Content-Type", content_type)
|
||||
self.send_header("Content-Length", str(len(data)))
|
||||
self.end_headers()
|
||||
self.wfile.write(data)
|
||||
|
||||
|
||||
def main():
|
||||
# 配置根日志
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
datefmt="%H:%M:%S",
|
||||
)
|
||||
|
||||
server = ThreadingHTTPServer((config.SERVER_HOST, config.SERVER_PORT), Handler)
|
||||
backend = config.LLM_BACKEND
|
||||
print("=" * 60)
|
||||
print(f"大模型速度测试助手 v2.1 ({'OpenAI 兼容' if backend == 'openai' else 'Ollama'})")
|
||||
print("=" * 60)
|
||||
if backend == "openai":
|
||||
print(f"服务地址 : {config.OPENAI_BASE_URL}")
|
||||
print(f"默认模型 : {config.OPENAI_DEFAULT_MODEL}")
|
||||
print(f"并发测试 : {config.DEFAULT_CONCURRENCY} 并发")
|
||||
else:
|
||||
print(f"Ollama 地址 : {config.OLLAMA_BASE_URL}")
|
||||
print(f"默认模型 : {config.DEFAULT_MODEL}")
|
||||
print(f"前端界面 : http://localhost:{config.SERVER_PORT}")
|
||||
print("按 Ctrl+C 停止服务器")
|
||||
print("=" * 60)
|
||||
try:
|
||||
server.serve_forever()
|
||||
except KeyboardInterrupt:
|
||||
print("\n服务器已停止")
|
||||
server.server_close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,61 @@
|
||||
"""统计计算工具"""
|
||||
import math
|
||||
from typing import List
|
||||
|
||||
|
||||
def _percentile_at(sorted_data: list[float], pct: float) -> float:
|
||||
"""在线性插值下,返回排序后数据的指定百分位值。
|
||||
|
||||
注意:调用方必须确保 data 已排序。
|
||||
"""
|
||||
n = len(sorted_data)
|
||||
if n == 0:
|
||||
return 0.0
|
||||
k = (pct / 100.0) * (n - 1)
|
||||
f = math.floor(k)
|
||||
c = math.ceil(k)
|
||||
if f == c:
|
||||
return sorted_data[int(k)]
|
||||
d0 = sorted_data[int(f)] * (c - k)
|
||||
d1 = sorted_data[int(c)] * (k - f)
|
||||
return d0 + d1
|
||||
|
||||
|
||||
def percentile(data: List[float], pct: float) -> float:
|
||||
"""计算指定百分位数的值(线性插值)
|
||||
|
||||
空数据返回 0.0,调用方应通过检查 count 来判断是否有数据。
|
||||
"""
|
||||
return _percentile_at(sorted(data), pct)
|
||||
|
||||
|
||||
def percentiles(data: List[float], pcts: List[float]) -> dict:
|
||||
"""单次排序,批量计算多个百分位,返回 {pct: value}。"""
|
||||
sorted_data = sorted(data)
|
||||
return {p: _percentile_at(sorted_data, p) for p in pcts}
|
||||
|
||||
|
||||
def mean(data: List[float]) -> float:
|
||||
"""算术平均值"""
|
||||
if not data:
|
||||
return 0.0
|
||||
return sum(data) / len(data)
|
||||
|
||||
|
||||
def std(data: List[float]) -> float:
|
||||
"""样本标准差(n-1)"""
|
||||
if len(data) < 2:
|
||||
return 0.0
|
||||
m = mean(data)
|
||||
variance = sum((x - m) ** 2 for x in data) / (len(data) - 1)
|
||||
return math.sqrt(variance)
|
||||
|
||||
|
||||
def min_val(data: List[float]) -> float:
|
||||
"""最小值"""
|
||||
return min(data) if data else 0.0
|
||||
|
||||
|
||||
def max_val(data: List[float]) -> float:
|
||||
"""最大值"""
|
||||
return max(data) if data else 0.0
|
||||
@@ -0,0 +1 @@
|
||||
"""测试用例包:每个用例实现 run_test(client, repeats, stop_event)"""
|
||||
@@ -0,0 +1,32 @@
|
||||
"""测试用例公共工具"""
|
||||
from ..ollama_client import InferenceResult
|
||||
|
||||
|
||||
def make_result(infer: InferenceResult, iteration: int, **extra) -> dict:
|
||||
"""将一次推理结果转换为用例数据条目"""
|
||||
item = {
|
||||
"iteration": iteration,
|
||||
"ttft_ms": round(infer.ttft_ms, 2),
|
||||
"prefill_ms": round(infer.prefill_ms, 2),
|
||||
"decode_speed_tok_s": round(infer.decode_speed_tok_s, 2),
|
||||
"total_tokens": infer.total_tokens,
|
||||
"prompt_tokens": infer.prompt_tokens,
|
||||
"completion_tokens": infer.completion_tokens,
|
||||
"e2e_ms": round(infer.e2e_ms, 2),
|
||||
"elapsed_ms": round(infer.e2e_ms, 2),
|
||||
"response_length": len(infer.response),
|
||||
}
|
||||
item.update(extra)
|
||||
return item
|
||||
|
||||
|
||||
def case_meta(case_id: str, name: str, description: str,
|
||||
difficulty: str, estimated_seconds: int) -> dict:
|
||||
"""构造用例元数据"""
|
||||
return {
|
||||
"case_id": case_id,
|
||||
"name": name,
|
||||
"description": description,
|
||||
"difficulty": difficulty,
|
||||
"estimated_seconds": estimated_seconds,
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
"""用例1:纯文本生成基准 —— 基础生成速度基准测试"""
|
||||
from ..ollama_client import OllamaClient
|
||||
from .base import case_meta, make_result
|
||||
|
||||
PROMPT = "请用200字左右概述人工智能从诞生至今的关键发展阶段,每段不超过50字。"
|
||||
|
||||
|
||||
def run_test(client: OllamaClient, repeats: int = 3, stop_event=None) -> dict:
|
||||
"""执行纯文本生成测试"""
|
||||
results = []
|
||||
for i in range(repeats):
|
||||
if stop_event and stop_event.is_set():
|
||||
break
|
||||
infer = client.generate(PROMPT, options={"temperature": 0.7, "num_predict": 300})
|
||||
results.append(make_result(infer, i + 1))
|
||||
|
||||
meta = case_meta(
|
||||
"case_01_generation", "纯文本生成基准",
|
||||
"基础生成速度基准测试", "简单", 5,
|
||||
)
|
||||
meta["results"] = results
|
||||
meta["prompt_used"] = PROMPT
|
||||
return meta
|
||||
@@ -0,0 +1,41 @@
|
||||
"""用例2:简单工具调用延迟 —— 测量结构化工具输出的调用延迟"""
|
||||
from ..ollama_client import OllamaClient
|
||||
from .base import case_meta, make_result
|
||||
|
||||
# 模拟"工具调用"场景:要求模型只输出结构化结果(短输出、低复杂度)
|
||||
SYSTEM = (
|
||||
"你是一个工具调用助手。你只能输出一个JSON对象,格式为 "
|
||||
'{"tool": "工具名", "params": {"参数名": "参数值"}},不要输出任何其它文字。'
|
||||
)
|
||||
|
||||
TOOL_PROMPTS = [
|
||||
("查询当前日期", {"tool": "get_date", "params": {}}),
|
||||
("获取CPU使用率", {"tool": "get_cpu_usage", "params": {}}),
|
||||
("列出当前目录文件", {"tool": "list_files", "params": {"path": "/tmp"}}),
|
||||
]
|
||||
|
||||
|
||||
def run_test(client: OllamaClient, repeats: int = 3, stop_event=None) -> dict:
|
||||
"""执行简单工具调用延迟测试"""
|
||||
results = []
|
||||
for i in range(repeats):
|
||||
if stop_event and stop_event.is_set():
|
||||
break
|
||||
desc, expected = TOOL_PROMPTS[i % len(TOOL_PROMPTS)]
|
||||
infer = client.generate(
|
||||
f"请调用工具:{desc}",
|
||||
system=SYSTEM,
|
||||
options={"temperature": 0.0, "num_predict": 60},
|
||||
)
|
||||
results.append(make_result(
|
||||
infer, i + 1,
|
||||
tool=expected["tool"],
|
||||
matched=(expected["tool"] in infer.response),
|
||||
))
|
||||
|
||||
meta = case_meta(
|
||||
"case_02_simple_tool", "简单工具调用延迟",
|
||||
"测量简单工具调用的额外延迟", "简单", 3,
|
||||
)
|
||||
meta["results"] = results
|
||||
return meta
|
||||
@@ -0,0 +1,72 @@
|
||||
"""用例3:文件读取+分析 —— 测试文件写入、读取、分析的完整链路"""
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
from ..ollama_client import OllamaClient
|
||||
from .base import case_meta, make_result
|
||||
|
||||
SAMPLE_CODE = '''"""示例模块:一个简单的计算器"""
|
||||
class Calculator:
|
||||
"""四则运算计算器"""
|
||||
|
||||
def __init__(self):
|
||||
self.history = []
|
||||
|
||||
def add(self, a, b):
|
||||
self.history.append(("add", a, b))
|
||||
return a + b
|
||||
|
||||
def divide(self, a, b):
|
||||
if b == 0:
|
||||
raise ValueError("除数不能为零")
|
||||
self.history.append(("divide", a, b))
|
||||
return a / b
|
||||
|
||||
|
||||
def main():
|
||||
calc = Calculator()
|
||||
print(calc.add(1, 2))
|
||||
print(calc.divide(10, 2))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
'''
|
||||
|
||||
ANALYSIS_PROMPT = "请分析以下Python代码的结构、功能,并指出其中可能存在的问题:\n\n```python\n{code}\n```"
|
||||
|
||||
|
||||
def run_test(client: OllamaClient, repeats: int = 3, stop_event=None) -> dict:
|
||||
"""执行文件读取+分析测试"""
|
||||
results = []
|
||||
for i in range(repeats):
|
||||
if stop_event and stop_event.is_set():
|
||||
break
|
||||
|
||||
# 写入临时文件,模拟"生成测试文件"环节
|
||||
fd, test_path = tempfile.mkstemp(prefix="llm_speed_test_", suffix=".py")
|
||||
try:
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as f:
|
||||
f.write(SAMPLE_CODE)
|
||||
|
||||
# 读取文件内容作为模型输入
|
||||
with open(test_path, "r", encoding="utf-8") as f:
|
||||
content = f.read()
|
||||
finally:
|
||||
try:
|
||||
os.remove(test_path)
|
||||
except OSError:
|
||||
pass
|
||||
|
||||
infer = client.generate(
|
||||
ANALYSIS_PROMPT.format(code=content),
|
||||
options={"temperature": 0.3, "num_predict": 300},
|
||||
)
|
||||
results.append(make_result(infer, i + 1, file_chars=len(content)))
|
||||
|
||||
meta = case_meta(
|
||||
"case_03_file_analysis", "文件读取+分析",
|
||||
"文件读取与分析完整链路计时", "中等", 2,
|
||||
)
|
||||
meta["results"] = results
|
||||
return meta
|
||||
@@ -0,0 +1,64 @@
|
||||
"""用例4:并行工具调用 —— 对比串行 vs 并行的加速比"""
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor
|
||||
|
||||
from ..ollama_client import OllamaClient
|
||||
from .base import case_meta
|
||||
|
||||
# 三个独立、互不依赖的"工具调用"请求
|
||||
PARALLEL_PROMPTS = [
|
||||
"请用一句话说明今天的天气如何。",
|
||||
"请用一句话说明Python的优缺点。",
|
||||
"请用一句话说明如何提高代码质量。",
|
||||
]
|
||||
|
||||
NUM_PARALLEL = len(PARALLEL_PROMPTS)
|
||||
|
||||
|
||||
def run_test(client: OllamaClient, repeats: int = 3, stop_event=None) -> dict:
|
||||
"""执行串行 vs 并行对比测试"""
|
||||
results = []
|
||||
|
||||
def call_one(prompt: str):
|
||||
return client.generate(prompt, options={"temperature": 0.3, "num_predict": 80})
|
||||
|
||||
# 线程池在循环外创建,复用 worker 线程
|
||||
with ThreadPoolExecutor(max_workers=NUM_PARALLEL) as executor:
|
||||
for i in range(repeats):
|
||||
if stop_event and stop_event.is_set():
|
||||
break
|
||||
|
||||
# 串行:依次执行
|
||||
serial_start = time.perf_counter()
|
||||
serial_results = [call_one(p) for p in PARALLEL_PROMPTS]
|
||||
serial_ms = (time.perf_counter() - serial_start) * 1000
|
||||
|
||||
# 并行:线程池并发执行
|
||||
parallel_start = time.perf_counter()
|
||||
parallel_results = list(executor.map(call_one, PARALLEL_PROMPTS))
|
||||
parallel_ms = (time.perf_counter() - parallel_start) * 1000
|
||||
|
||||
speedup = (serial_ms / parallel_ms) if parallel_ms > 0 else 0.0
|
||||
|
||||
avg_ttft = sum(r.ttft_ms for r in parallel_results) / NUM_PARALLEL
|
||||
avg_decode = sum(r.decode_speed_tok_s for r in parallel_results) / NUM_PARALLEL
|
||||
|
||||
results.append({
|
||||
"iteration": i + 1,
|
||||
"parallel_tasks": NUM_PARALLEL,
|
||||
"serial_ms": round(serial_ms, 2),
|
||||
"parallel_ms": round(parallel_ms, 2),
|
||||
"speedup_ratio": round(speedup, 2),
|
||||
"elapsed_ms": round(parallel_ms, 2), # 用并行耗时作为该用例的代表延迟
|
||||
"avg_ttft_ms": round(avg_ttft, 2),
|
||||
"avg_decode_speed_tok_s": round(avg_decode, 2),
|
||||
"ttft_ms": round(avg_ttft, 2),
|
||||
"decode_speed_tok_s": round(avg_decode, 2),
|
||||
})
|
||||
|
||||
meta = case_meta(
|
||||
"case_04_parallel_tool", "并行工具调用",
|
||||
"对比串行 vs 并行工具调用的延迟差异", "中等", 5,
|
||||
)
|
||||
meta["results"] = results
|
||||
return meta
|
||||
@@ -0,0 +1,45 @@
|
||||
"""用例5:长上下文处理 —— 测试不同上下文大小下的处理延迟"""
|
||||
from ..ollama_client import OllamaClient
|
||||
from .base import case_meta, make_result
|
||||
|
||||
# 约 15 个 token/句的填充句,用于按目标 token 数拼接上下文
|
||||
FILLER_SENTENCE = "这是一个用于测试长上下文处理能力的示例句子,包含一些常见的词汇和表述方式。\n"
|
||||
TOKENS_PER_SENTENCE = 15
|
||||
|
||||
CONTEXT_SIZES = [500, 1000, 2000, 4000] # 目标 prompt token 数
|
||||
|
||||
QUESTION = "请用一句话总结:上面这段文本主要介绍了什么?"
|
||||
|
||||
|
||||
def _build_context(target_tokens: int) -> str:
|
||||
"""按目标 token 数近似生成一段填充文本"""
|
||||
n_sentences = max(1, target_tokens // TOKENS_PER_SENTENCE)
|
||||
return FILLER_SENTENCE * n_sentences
|
||||
|
||||
|
||||
def run_test(client: OllamaClient, repeats: int = 2, stop_event=None) -> dict:
|
||||
"""执行长上下文处理测试"""
|
||||
results = []
|
||||
for context_size in CONTEXT_SIZES:
|
||||
for r in range(repeats):
|
||||
if stop_event and stop_event.is_set():
|
||||
break
|
||||
context = _build_context(context_size)
|
||||
prompt = f"{context}\n\n{QUESTION}"
|
||||
infer = client.generate(
|
||||
prompt,
|
||||
options={"temperature": 0.3, "num_predict": 80},
|
||||
)
|
||||
results.append(make_result(
|
||||
infer, r + 1,
|
||||
context_tokens_target=context_size,
|
||||
context_chars=len(context),
|
||||
))
|
||||
|
||||
meta = case_meta(
|
||||
"case_05_long_context", "长上下文处理",
|
||||
"不同上下文长度对延迟的影响", "耗时", 10,
|
||||
)
|
||||
meta["results"] = results
|
||||
meta["context_sizes_tested"] = CONTEXT_SIZES
|
||||
return meta
|
||||
@@ -0,0 +1,33 @@
|
||||
"""用例6:复杂推理任务 —— 对比不同难度推理任务的延迟差异"""
|
||||
from ..ollama_client import OllamaClient
|
||||
from .base import case_meta, make_result
|
||||
|
||||
REASONING_LEVELS = [
|
||||
("简单问答", "什么是人工智能?", 80),
|
||||
("中等推理", "解释Transformer架构中的自注意力机制是如何工作的。", 200),
|
||||
("复杂推理", "分析大语言模型在长链推理中的主要局限性,并提出三种改进方案。", 400),
|
||||
]
|
||||
|
||||
|
||||
def run_test(client: OllamaClient, repeats: int = 2, stop_event=None) -> dict:
|
||||
"""执行复杂推理任务测试"""
|
||||
results = []
|
||||
for level_name, prompt, num_predict in REASONING_LEVELS:
|
||||
for i in range(repeats):
|
||||
if stop_event and stop_event.is_set():
|
||||
break
|
||||
infer = client.generate(
|
||||
prompt,
|
||||
options={"temperature": 0.3, "num_predict": num_predict},
|
||||
)
|
||||
results.append(make_result(
|
||||
infer, i + 1,
|
||||
reasoning_level=level_name,
|
||||
))
|
||||
|
||||
meta = case_meta(
|
||||
"case_06_reasoning", "复杂推理任务",
|
||||
"不同推理复杂度下的延迟对比", "耗时", 15,
|
||||
)
|
||||
meta["results"] = results
|
||||
return meta
|
||||
@@ -0,0 +1,49 @@
|
||||
"""用例7:多轮对话累积 —— 追踪连续多轮对话中延迟随上下文累积的变化"""
|
||||
from ..ollama_client import OllamaClient
|
||||
from .base import case_meta, make_result
|
||||
|
||||
DEBUG_SCENARIOS = [
|
||||
"修复Python中的索引越界错误",
|
||||
"修复JavaScript闭包变量捕获问题",
|
||||
"修复SQL注入漏洞",
|
||||
"修复多线程竞态条件",
|
||||
"修复递归栈溢出",
|
||||
"修复正则表达式灾难回溯",
|
||||
]
|
||||
|
||||
|
||||
def run_test(client: OllamaClient, repeats: int = 1, stop_event=None) -> dict:
|
||||
"""执行多轮对话累积测试(使用 /api/chat 累积上下文)
|
||||
|
||||
repeats 控制总轮次数,每轮切换不同调试场景。
|
||||
"""
|
||||
results = []
|
||||
messages = []
|
||||
|
||||
for i in range(repeats):
|
||||
if stop_event and stop_event.is_set():
|
||||
break
|
||||
scenario = DEBUG_SCENARIOS[i % len(DEBUG_SCENARIOS)]
|
||||
|
||||
user_msg = f"(对话第{i + 1}轮)请协助:{scenario}。请简要给出你的分析。"
|
||||
messages.append({"role": "user", "content": user_msg})
|
||||
|
||||
infer = client.chat(messages, options={"temperature": 0.3, "num_predict": 120})
|
||||
|
||||
assistant_reply = infer.response
|
||||
messages.append({"role": "assistant", "content": assistant_reply})
|
||||
|
||||
results.append(make_result(
|
||||
infer, i + 1,
|
||||
round=i + 1,
|
||||
scenario=scenario,
|
||||
messages_in_context=len(messages),
|
||||
))
|
||||
|
||||
meta = case_meta(
|
||||
"case_07_multiturn", "多轮对话累积",
|
||||
"追踪连续多轮对话中延迟随上下文累积的变化", "耗时", 10,
|
||||
)
|
||||
meta["results"] = results
|
||||
meta["total_rounds"] = len(results)
|
||||
return meta
|
||||
@@ -0,0 +1,146 @@
|
||||
"""用例8:并发压力测试 —— 测量高并发下的 TPS 保持率、P95/P99 延迟与错误率
|
||||
|
||||
并发数取 config.DEFAULT_CONCURRENCY(默认 100),repeats 控制并发批次数量。
|
||||
"""
|
||||
import concurrent.futures
|
||||
import logging
|
||||
|
||||
from .. import config
|
||||
from ..stats import percentile
|
||||
from .base import case_meta, make_result
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
CONCURRENT_PROMPT = "用一句话回答:1+1等于几?"
|
||||
BASELINE_SAMPLES = 5 # 基准串行采样次数
|
||||
|
||||
|
||||
def _run_one(client) -> dict:
|
||||
"""单次并发请求,失败时返回 None 标记"""
|
||||
try:
|
||||
infer = client.generate(
|
||||
CONCURRENT_PROMPT,
|
||||
options={"temperature": 0.0, "num_predict": 32},
|
||||
)
|
||||
return {
|
||||
"ok": True,
|
||||
"infer": infer,
|
||||
"elapsed_ms": infer.e2e_ms,
|
||||
}
|
||||
except Exception as e:
|
||||
logger.debug("并发请求失败: %s", e)
|
||||
return {"ok": False, "error": str(e)}
|
||||
|
||||
|
||||
def _collect_serial(client, n) -> list:
|
||||
"""串行基准采样"""
|
||||
entries = []
|
||||
for i in range(n):
|
||||
r = _run_one(client)
|
||||
if r["ok"]:
|
||||
entries.append(r)
|
||||
return entries
|
||||
|
||||
|
||||
def _collect_concurrent(client, n, stop_event=None) -> list:
|
||||
"""并发采集(线程池)"""
|
||||
entries = []
|
||||
completed = 0
|
||||
|
||||
def _wrapped():
|
||||
return _run_one(client)
|
||||
|
||||
with concurrent.futures.ThreadPoolExecutor(max_workers=n) as executor:
|
||||
futures = [executor.submit(_wrapped) for _ in range(n)]
|
||||
for fut in concurrent.futures.as_completed(futures):
|
||||
try:
|
||||
r = fut.result(timeout=300)
|
||||
except Exception as e:
|
||||
r = {"ok": False, "error": str(e)}
|
||||
entries.append(r)
|
||||
completed += 1
|
||||
if stop_event and stop_event.is_set():
|
||||
for f in futures:
|
||||
f.cancel()
|
||||
break
|
||||
|
||||
return entries
|
||||
|
||||
|
||||
def run_test(client, repeats: int = 1, stop_event=None) -> dict:
|
||||
"""执行并发压力测试"""
|
||||
if repeats > 3:
|
||||
repeats = 3 # 并发批次过多会非常耗时,限制上限
|
||||
|
||||
n_concurrent = config.DEFAULT_CONCURRENCY
|
||||
|
||||
# ---- Step 1: 串行基准 ----
|
||||
logger.info("[并发] 计算基准 TPS (串行 %d 次)...", BASELINE_SAMPLES)
|
||||
baseline = _collect_serial(client, BASELINE_SAMPLES)
|
||||
valid_base = [e for e in baseline if e["ok"]]
|
||||
if not valid_base:
|
||||
return case_meta(
|
||||
"case_08_concurrency", "并发压力测试",
|
||||
"基准测试全部失败", "压力", 30,
|
||||
) | {"status": "failed", "error": "基准测试全部失败", "results": []}
|
||||
|
||||
base_tps = sum(e["infer"].decode_speed_tok_s for e in valid_base) / len(valid_base)
|
||||
base_e2e = sorted(e["elapsed_ms"] for e in valid_base)
|
||||
base_p50 = percentile(base_e2e, 50)
|
||||
logger.info("[并发] 基准 TPS: %.2f tok/s, E2E P50: %.0fms", base_tps, base_p50)
|
||||
|
||||
# ---- Step 2: 并发测试 ----
|
||||
logger.info("[并发] 启动 %d 并发请求...", n_concurrent)
|
||||
all_concurrent = []
|
||||
for batch in range(repeats):
|
||||
if stop_event and stop_event.is_set():
|
||||
break
|
||||
batch_results = _collect_concurrent(client, n_concurrent, stop_event)
|
||||
all_concurrent.extend(batch_results)
|
||||
ok_count = sum(1 for r in batch_results if r["ok"])
|
||||
logger.info("[并发] 批次 %d 完成: %d/%d 成功", batch + 1, ok_count, len(batch_results))
|
||||
|
||||
# ---- Step 3: 汇总 ----
|
||||
valid_con = [e for e in all_concurrent if e["ok"]]
|
||||
failed_count = len(all_concurrent) - len(valid_con)
|
||||
|
||||
con_tps = (
|
||||
sum(e["infer"].decode_speed_tok_s for e in valid_con) / len(valid_con)
|
||||
if valid_con else 0.0
|
||||
)
|
||||
e2e_vals = sorted(e["elapsed_ms"] for e in valid_con)
|
||||
e2e_p50 = percentile(e2e_vals, 50)
|
||||
e2e_p95 = percentile(e2e_vals, 95)
|
||||
e2e_p99 = percentile(e2e_vals, 99)
|
||||
|
||||
tps_keep = (con_tps / base_tps * 100) if base_tps > 0 else 0.0
|
||||
total = len(all_concurrent) or 1
|
||||
error_rate = failed_count / total * 100
|
||||
|
||||
# 构造标准结果条目(供引擎统计/前端展示),并附加并发摘要
|
||||
results = [
|
||||
make_result(
|
||||
e["infer"], i + 1,
|
||||
batch="concurrent",
|
||||
error=("" if e["ok"] else e.get("error", "unknown")),
|
||||
)
|
||||
for i, e in enumerate(all_concurrent)
|
||||
if e["ok"]
|
||||
]
|
||||
|
||||
meta = case_meta(
|
||||
"case_08_concurrency", "并发压力测试",
|
||||
"高并发下的 TPS 保持率、P95/P99 延迟、错误率", "压力", 30,
|
||||
)
|
||||
meta["results"] = results
|
||||
meta["concurrency"] = n_concurrent
|
||||
meta["baseline_tps"] = round(base_tps, 2)
|
||||
meta["concurrent_tps"] = round(con_tps, 2)
|
||||
meta["tps_keep_rate_pct"] = round(tps_keep, 2)
|
||||
meta["e2e_p50_ms"] = round(e2e_p50, 2)
|
||||
meta["e2e_p95_ms"] = round(e2e_p95, 2)
|
||||
meta["e2e_p99_ms"] = round(e2e_p99, 2)
|
||||
meta["success_count"] = len(valid_con)
|
||||
meta["failed_count"] = failed_count
|
||||
meta["error_rate_pct"] = round(error_rate, 2)
|
||||
return meta
|
||||
Reference in New Issue
Block a user