后端 (platform.py + platform_store.py): - 新增 _build_messages_payload() 转换前端格式为 OpenAI messages - 新增 _stream_chat_proxy() SSE 流式代理到算力节点 - 新增 _unload_from_compute_node() 真正释放算力节点 GPU 显存 - 重写 model_compare_load: 从假 PID/端口改为真正调用算力节点加载模型 - 修复 model_compare_unload: 调用 _unload_from_compute_node 释放 GPU - 修复 model_compare_delete: 先释放 GPU 再删除记录 - 修复 model_compare_stream_chat: 从 mock 改为 StreamingResponse 代理 - 修复 model_chat_local/stream: 消息格式转换 + 路径修正 - PlatformStore 新增 _inference_nodes 追踪,gpus() 同步推理占用状态 - preload/unload 端点标记/清除推理节点占用 算力节点 (compute): - inference.py: 适配新版 LLaMA-Factory API (get_infer_args 4 返回值、ChatModel args dict、stream_chat 新签名) - inference.py: unload() 增加 gc.collect + torch.cuda.empty_cache + synchronize 彻底释放显存 - main.py: inference/load 移除 HTTPException(500),错误以 200 正常返回 前端: - InferenceChatView: 真实模式下走 SSE 流式推理,mock 模式保留兼容 - InferenceCreateView: 调用 preloadLocalModel + createCompare 真实创建推理任务,失败回退 mock - InferenceListView: 「停止」改为「释放」,删除前先释放算力节点,改进错误提示 - compare.ts: 新增 streamChatReal() fetch SSE,preload 超时提升至 5 分钟 - useStreamChat.ts: send() 支持 useMock 参数,真实模式调用 streamChatReal - GPU 选择过滤: 仅显示在线算力节点上的空闲 GPU Co-Authored-By: Claude <noreply@anthropic.com>
146 lines
5.9 KiB
Python
146 lines
5.9 KiB
Python
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
|