Initial commit: LocalPilot:本地模型运行时,复用 llama-server 并提供 Ollama 兼容 provider

This commit is contained in:
WpyQwq
2026-09-19 11:57:59 +08:00
commit d159212416
27 changed files with 3022 additions and 0 deletions
+178
View File
@@ -0,0 +1,178 @@
from __future__ import annotations
import asyncio
import queue
import threading
from collections.abc import AsyncIterator
from pathlib import Path
from typing import Any
from ..config import ModelConfig, RuntimeConfig
from ..types import BackendError, ChatChunk, ChatParams
from .base import BaseBackend
class ONNXBackend(BaseBackend):
"""ONNX loader with an Optimum path when installed, plus safe inspection fallback."""
def __init__(self, model: ModelConfig, runtime: RuntimeConfig) -> None:
super().__init__(model, runtime)
self.session: Any = None
self.ort_model: Any = None
self.tokenizer: Any = None
self.providers: list[str] = []
async def load(self) -> None:
if self.loaded:
return
if not self.model.path:
raise BackendError(f"ONNX 模型缺少 path: {self.model.id}")
path = Path(self.model.path)
if not path.exists():
raise BackendError(f"ONNX 文件或目录不存在: {path}")
try:
import onnxruntime as ort
except ImportError as exc:
raise BackendError("当前 Conda LLM 环境没有 onnxruntime") from exc
providers = ["CUDAExecutionProvider", "CPUExecutionProvider"]
available = ort.get_available_providers()
providers = [provider for provider in providers if provider in available]
if not providers:
providers = available
self.providers = providers
try:
if path.is_file():
self.session = await asyncio.to_thread(ort.InferenceSession, str(path), providers=providers)
else:
onnx_file = next(path.glob("*.onnx"), None)
if onnx_file is None:
raise BackendError(f"目录中没有 .onnx 文件: {path}")
self.session = await asyncio.to_thread(
ort.InferenceSession, str(onnx_file), providers=providers
)
except Exception as exc:
raise BackendError(f"ONNX Runtime 加载失败: {exc}") from exc
try:
from optimum.onnxruntime import ORTModelForCausalLM
from transformers import AutoTokenizer
model_dir = path if path.is_dir() else path.parent
self.ort_model = await asyncio.to_thread(
ORTModelForCausalLM.from_pretrained,
model_dir,
provider=providers[0] if providers else "CPUExecutionProvider",
)
self.tokenizer = await asyncio.to_thread(
AutoTokenizer.from_pretrained, model_dir, local_files_only=True
)
except ImportError:
self.ort_model = None
self.loaded = True
async def unload(self) -> None:
self.session = None
self.ort_model = None
self.tokenizer = None
self.loaded = False
async def chat(
self,
messages: list[dict[str, Any]],
params: ChatParams,
stream: bool = False,
) -> AsyncIterator[ChatChunk]:
if not self.loaded:
await self.load()
if self.ort_model is None or self.tokenizer is None:
raise BackendError(
"ONNX 文件已加载并可检查,但聊天生成需要安装 optimum[onnxruntime],"
"并使用带 tokenizer/config 的 ONNX CausalLM 目录。"
)
kwargs: dict[str, Any] = {
"add_generation_prompt": True,
"tokenize": True,
"return_tensors": "pt",
"return_dict": True,
}
if params.enable_thinking is not None:
kwargs["enable_thinking"] = params.enable_thinking
try:
inputs = self.tokenizer.apply_chat_template(messages, **kwargs)
except TypeError:
kwargs.pop("enable_thinking", None)
inputs = self.tokenizer.apply_chat_template(messages, **kwargs)
generation_kwargs: dict[str, Any] = {
"max_new_tokens": params.max_tokens,
"do_sample": params.temperature > 0,
"use_cache": True,
"repetition_penalty": params.repeat_penalty,
"pad_token_id": self.tokenizer.pad_token_id,
"eos_token_id": self.tokenizer.eos_token_id,
}
if params.temperature > 0:
generation_kwargs.update({"temperature": params.temperature, "top_p": params.top_p, "top_k": params.top_k})
generation_kwargs.update(params.extra)
if not stream:
output = await asyncio.to_thread(self.ort_model.generate, **inputs, **generation_kwargs)
prompt_len = int(inputs["input_ids"].shape[-1])
text = self.tokenizer.decode(output[0][prompt_len:], skip_special_tokens=True)
for stop in params.stop:
text = text.split(stop, 1)[0]
yield ChatChunk(text=text, done=True, finish_reason="stop")
return
from transformers import TextIteratorStreamer
streamer = TextIteratorStreamer(self.tokenizer, skip_prompt=True, skip_special_tokens=True)
events: queue.Queue[tuple[str, Any]] = queue.Queue()
def worker() -> None:
try:
self.ort_model.generate(**inputs, streamer=streamer, **generation_kwargs)
except BaseException as exc:
events.put(("error", exc))
try:
streamer.end()
except Exception:
pass
def forward_stream() -> None:
try:
for item in streamer:
events.put(("text", item))
events.put(("done", None))
except BaseException as exc:
events.put(("error", exc))
threading.Thread(target=worker, daemon=True).start()
threading.Thread(target=forward_stream, daemon=True).start()
while True:
kind, value = await asyncio.to_thread(events.get)
if kind == "text":
text = str(value)
for stop in params.stop:
text = text.split(stop, 1)[0]
if text:
yield ChatChunk(text=text)
elif kind == "error":
raise BackendError(f"ONNX 推理失败: {value}") from value
else:
yield ChatChunk(done=True, finish_reason="stop")
return
def info(self) -> dict[str, Any]:
inputs: list[str] = []
if self.session is not None:
inputs = [item.name for item in self.session.get_inputs()]
return {
"id": self.model.id,
"kind": "onnx",
"path": self.model.path,
"loaded": self.loaded,
"providers": self.providers,
"inputs": inputs,
"generation_ready": self.ort_model is not None and self.tokenizer is not None,
}