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
+154
View File
@@ -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