Files
llm_speed_test_app/backend/ollama_client.py
T

155 lines
5.7 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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