Files
llm_speed_test_app/backend/test_cases/case_08_concurrency.py
T

147 lines
4.9 KiB
Python

"""用例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