from __future__ import annotations import threading import time from typing import Any class InferenceSession: """Manages a loaded model for inference with LLaMA-Factory ChatModel.""" def __init__(self) -> None: self._model: Any = None self._tokenizer: Any = None self._generating_args: dict[str, Any] = {} self._model_name: str = "" self._adapter_path: str = "" self._lock = threading.Lock() self._loaded_at: float = 0.0 self._status: str = "idle" @property def status(self) -> str: return self._status @property def model_name(self) -> str: return self._model_name @property def adapter_path(self) -> str: return self._adapter_path @property def loaded_at(self) -> float: return self._loaded_at def info(self) -> dict[str, Any]: return { "loaded": self._status == "ready", "status": self._status, "model_name": self._model_name, "adapter_path": self._adapter_path, "loaded_at": self._loaded_at, } def load(self, model_name_or_path, adapter_name_or_path="", template="qwen", infer_backend="huggingface", infer_dtype="auto", **kwargs): with self._lock: if self._status == "loading": return {"loaded": False, "error": "model is already loading"} if self._status == "ready": self.unload() self._status = "loading" self._model_name = model_name_or_path self._adapter_path = adapter_name_or_path try: from llamafactory.chat import ChatModel from llamafactory.hparams import get_infer_args args = {"model_name_or_path": model_name_or_path, "template": template, "infer_backend": infer_backend, "infer_dtype": infer_dtype} if adapter_name_or_path: args["adapter_name_or_path"] = adapter_name_or_path args.update(kwargs) infer_result = get_infer_args(args) # ChatModel internally re-parses the args dict via get_infer_args, # so pass the original args (not the parsed dataclass objects). self._model = ChatModel(args) self._tokenizer = getattr(self._model, 'tokenizer', None) or self._model.engine.tokenizer # Extract generating_args (last element) for later use in chat() generating_args = infer_result[-1] if hasattr(generating_args, '__dataclass_fields__'): self._generating_args = {k: v for k, v in vars(generating_args).items() if not k.startswith('_')} else: self._generating_args = dict(generating_args) self._loaded_at = time.time() self._status = "ready" return {"loaded": True, "status": "ready"} except Exception as exc: self._status = "error" self._model = None return {"loaded": False, "status": "error", "error": str(exc)} def unload(self): with self._lock: if self._model is not None: try: del self._model except Exception: pass self._model = None self._tokenizer = None # 强制释放 PyTorch CUDA 缓存,真正归还 GPU 显存 try: import gc gc.collect() import torch if torch.cuda.is_available(): torch.cuda.empty_cache() torch.cuda.synchronize() except Exception: pass self._status = "idle" self._model_name = "" self._adapter_path = "" self._loaded_at = 0.0 return {"unloaded": True} def chat(self, messages, temperature=0.95, top_p=0.7, max_new_tokens=1024, do_sample=True, **kwargs): with self._lock: if self._status != "ready" or self._model is None: return {"error": "model not loaded", "response": ""} try: generate_kwargs = {"temperature": temperature, "top_p": top_p, "max_new_tokens": max_new_tokens, "do_sample": do_sample} generate_kwargs.update(kwargs) system = next((m["content"] for m in messages if m["role"] == "system"), None) user_messages = [m for m in messages if m["role"] != "system"] responses = [] for response in self._model.stream_chat(user_messages, system=system, **generate_kwargs): responses.append(response) full_response = "".join(str(r) for r in responses) return {"response": full_response} except Exception as exc: return {"error": str(exc), "response": ""} def chat_stream(self, messages, **kwargs): with self._lock: if self._status != "ready" or self._model is None: yield 'data: {"error": "model not loaded"}\n\n' return try: generate_kwargs = {**kwargs} system = next((m["content"] for m in messages if m["role"] == "system"), None) user_messages = [m for m in messages if m["role"] != "system"] for new_text in self._model.stream_chat(user_messages, system=system, **generate_kwargs): yield new_text except Exception as exc: yield 'data: {"error": "' + str(exc) + '"}\n\n' _inference_session = None def get_inference_session(): global _inference_session if _inference_session is None: _inference_session = InferenceSession() return _inference_session