Initial commit: LocalPilot:本地模型运行时,复用 llama-server 并提供 Ollama 兼容 provider
This commit is contained in:
@@ -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,
|
||||
}
|
||||
Reference in New Issue
Block a user