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
+1
View File
@@ -0,0 +1 @@
"""测试用例包:每个用例实现 run_test(client, repeats, stop_event)"""
+32
View File
@@ -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,
}
+23
View File
@@ -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
+41
View File
@@ -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
+33
View File
@@ -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
+49
View File
@@ -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
+146
View File
@@ -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