179 lines
6.8 KiB
Python
179 lines
6.8 KiB
Python
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,
|
||
}
|