This commit is contained in:
wangjiming
2026-07-31 16:10:34 +08:00
parent 945b4ace86
commit 242407b676
34 changed files with 3847 additions and 717 deletions

View File

@@ -0,0 +1,125 @@
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)
model_args, generating_args = get_infer_args(args)
self._model = ChatModel(model_args)
self._tokenizer = self._model.tokenizer
self._generating_args = 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
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 = {**self._generating_args, "temperature": temperature, "top_p": top_p, "max_new_tokens": max_new_tokens, "do_sample": do_sample}
generate_kwargs.update(kwargs)
formatted = self._model.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
responses = []
for response in self._model.stream_chat(formatted, 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 = {**self._generating_args, **kwargs}
formatted = self._model.tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
for new_text in self._model.stream_chat(formatted, 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