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, }