feat: 基于 LLaMA-Factory 实现模型推理引擎

利用容器内已有的 LLaMA-Factory ChatModel (huggingface 后端) 实现真实
模型推理,无需额外安装 vLLM。

Compute 端新增:
- compute/engines/llama_factory/inference.py
  InferenceSession: 模型加载/卸载/对话/流式输出
  支持 base model 和 LoRA adapter,线程安全
- compute/api/main.py 新增 5 个推理端点:
  POST /inference/load      - 加载模型
  POST /inference/unload    - 卸载释放 GPU 显存
  GET  /inference/status    - 查询会话状态
  POST /inference/chat      - 非流式对话
  POST /inference/chat/stream - SSE 流式对话

后端新增:
- platform.py 推理代理端点(local/chat, local/chat/stream,
  local/preload, local/unload, local/status, trained/preload)
- ComputeNodeClient._request 通用请求方法
- _select_first_online_node 自动选择在线节点

Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
wuyongtao
2026-07-28 13:49:10 +08:00
parent a9ab130d43
commit f917a025e1
4 changed files with 314 additions and 5 deletions

View File

@@ -10,10 +10,11 @@ from pathlib import Path
from typing import Any
from fastapi import FastAPI, File, Form, HTTPException, Query, Request, UploadFile
from fastapi.responses import FileResponse, JSONResponse
from fastapi.responses import FileResponse, JSONResponse, StreamingResponse
from compute.agent.process_manager import ProcessManager
from compute.engines.llama_factory.adapter import build_command, parse_log_line, prepare_runtime_files
from compute.engines.llama_factory.inference import get_inference_session
def create_app() -> FastAPI:
@@ -666,6 +667,85 @@ def create_app() -> FastAPI:
metrics = [parse_log_line(line) for line in window["content"].splitlines()]
return {"job_id": job_id, **window, "metrics": [m for m in metrics if m]}
# ── Inference Endpoints ───────────────────────────────────────────
@app.post(f"{route_prefix}/inference/load")
async def inference_load(payload: dict[str, Any]) -> dict[str, Any]:
"""Load a model for inference using LLaMA-Factory ChatModel.
Expected payload:
model_name_or_path: str (required)
adapter_name_or_path: str (optional, for LoRA adapters)
template: str (default: "qwen")
infer_backend: str (default: "huggingface")
infer_dtype: str (default: "auto")
"""
session = get_inference_session()
result = session.load(
model_name_or_path=payload.get("model_name_or_path", ""),
adapter_name_or_path=payload.get("adapter_name_or_path", ""),
template=payload.get("template", "qwen"),
infer_backend=payload.get("infer_backend", "huggingface"),
infer_dtype=payload.get("infer_dtype", "auto"),
)
if not result.get("loaded"):
raise HTTPException(status_code=500, detail=result.get("error", "model load failed"))
return result
@app.post(f"{route_prefix}/inference/unload")
async def inference_unload() -> dict[str, Any]:
"""Unload the currently loaded model and free GPU memory."""
return get_inference_session().unload()
@app.get(f"{route_prefix}/inference/status")
async def inference_status() -> dict[str, Any]:
"""Get the current inference session status."""
return get_inference_session().info()
@app.post(f"{route_prefix}/inference/chat")
async def inference_chat(payload: dict[str, Any]) -> dict[str, Any]:
"""Chat with the loaded model (non-streaming).
Expected payload:
messages: list[dict] (OpenAI format)
temperature: float (default 0.95)
top_p: float (default 0.7)
max_new_tokens: int (default 1024)
"""
messages = payload.get("messages") or []
if not messages:
raise HTTPException(status_code=400, detail="messages is required")
result = get_inference_session().chat(
messages=messages,
temperature=float(payload.get("temperature", 0.95)),
top_p=float(payload.get("top_p", 0.7)),
max_new_tokens=int(payload.get("max_new_tokens", 1024)),
do_sample=bool(payload.get("do_sample", True)),
)
if result.get("error"):
raise HTTPException(status_code=500, detail=result["error"])
return {"response": result["response"]}
@app.post(f"{route_prefix}/inference/chat/stream")
async def inference_chat_stream(payload: dict[str, Any]) -> StreamingResponse:
"""Chat with streaming response (Server-Sent Events)."""
messages = payload.get("messages") or []
if not messages:
raise HTTPException(status_code=400, detail="messages is required")
def generate():
session = get_inference_session()
for chunk in session.chat_stream(
messages=messages,
temperature=float(payload.get("temperature", 0.95)),
top_p=float(payload.get("top_p", 0.7)),
max_new_tokens=int(payload.get("max_new_tokens", 1024)),
do_sample=bool(payload.get("do_sample", True)),
):
yield chunk
return StreamingResponse(generate(), media_type="text/event-stream")
@app.post(f"{route_prefix}/compute/files/upload")
async def upload_file(
file: UploadFile | None = File(default=None),

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