feat: 模型推理异步加载与对话链路修复,同步基线
模型推理全异步化改造: - 计算节点 InferenceSession 改为后台线程异步加载模型,load 立即返回, 加载期间事件循环保持响应(/inference/status 与 /health 不阻塞) - 后端模型加载改为异步派发 + 轮询对账器(reconcile_inference_loads), 任务状态由 starting 自动推进到 ready/error,解决多节点启动超时 (timeout of 120000ms exceeded) - 推理删除/卸载改为任务感知 + 短超时,删除先删记录再 best-effort 卸载, 不再被不可达节点阻塞;同节点新模型替换旧任务标记失效 - 流式对话透传 task_id/node_id 路由到真正加载模型的算力节点, useStreamChat 解析 SSE 错误帧以干净文案展示 - 对话历史按任务 id 本地持久化,退出重进可恢复;移除页脚提示文本 - 新增后端推理异步加载与计算节点异步状态机单元测试 Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -33,6 +33,15 @@ def _select_first_online_node(store: Any) -> dict[str, Any] | None:
|
||||
return None
|
||||
|
||||
|
||||
def _candidate_online_nodes(store: Any, preferred_node_id: str | None = None) -> list[dict[str, Any]]:
|
||||
nodes = [node for node in store.compute_nodes() if node.get("enabled") and node.get("scheduler_status") == "online"]
|
||||
if not preferred_node_id:
|
||||
return nodes
|
||||
preferred = [node for node in nodes if node.get("id") == preferred_node_id]
|
||||
others = [node for node in nodes if node.get("id") != preferred_node_id]
|
||||
return preferred + others
|
||||
|
||||
|
||||
def _build_messages_payload(payload: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Convert frontend inference payload to compute API messages format.
|
||||
|
||||
@@ -61,10 +70,47 @@ def _build_messages_payload(payload: dict[str, Any]) -> dict[str, Any]:
|
||||
}
|
||||
|
||||
|
||||
def _node_for_inference_payload(store: Any, payload: dict[str, Any]) -> dict[str, Any] | None:
|
||||
node_id = payload.get("node_id") or payload.get("compute_node_id")
|
||||
task_id = payload.get("task_id") or payload.get("compare_task_id")
|
||||
if task_id and not node_id:
|
||||
try:
|
||||
task = store.compare_task(str(task_id))
|
||||
load_status = task.get("load_status") or {}
|
||||
if isinstance(load_status, str):
|
||||
load_status = json.loads(load_status)
|
||||
loaded_models = load_status.get("loaded_models") or []
|
||||
ready_model = next((item for item in loaded_models if item.get("status") in {"ready", "running"} and item.get("node_id")), None)
|
||||
if ready_model:
|
||||
node_id = ready_model.get("node_id")
|
||||
except Exception:
|
||||
node_id = None
|
||||
if node_id:
|
||||
return next((node for node in store.compute_nodes() if node.get("id") == node_id), None)
|
||||
return _select_first_online_node(store)
|
||||
|
||||
|
||||
async def _stream_chat_proxy(payload: dict[str, Any]) -> StreamingResponse:
|
||||
"""Common SSE streaming proxy: convert payload → forward to compute node → stream back."""
|
||||
store = get_platform_store()
|
||||
node = _select_first_online_node(store)
|
||||
# 任务仍在加载中时,直接返回明确的加载中提示,避免转发到尚未就绪的节点
|
||||
task_id = payload.get("task_id") or payload.get("compare_task_id")
|
||||
if task_id:
|
||||
try:
|
||||
task = store.compare_task(str(task_id))
|
||||
load_status = task.get("load_status") or {}
|
||||
if isinstance(load_status, str):
|
||||
load_status = json.loads(load_status)
|
||||
items = load_status.get("loaded_models") or []
|
||||
if items and not any(item.get("status") in {"ready", "running"} for item in items):
|
||||
if any(item.get("status") == "starting" for item in items):
|
||||
return StreamingResponse(
|
||||
iter(['data: {"error": "模型加载中,请稍候再试"}\n\n']),
|
||||
media_type="text/event-stream",
|
||||
)
|
||||
except Exception: # noqa: BLE001 - fall through to normal routing on lookup errors
|
||||
pass
|
||||
node = _node_for_inference_payload(store, payload)
|
||||
if not node:
|
||||
return StreamingResponse(
|
||||
iter(['data: {"error": "no online compute node available for inference"}\n\n']),
|
||||
@@ -729,12 +775,17 @@ async def merge_model(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
None,
|
||||
)
|
||||
base_model_path = payload.get("base_model_path") or (trained_model and trained_model.get("base_model_path"))
|
||||
adapter_path = payload.get("adapter_path") or payload.get("adapter_name_or_path") or (trained_model and trained_model.get("merged_path"))
|
||||
adapter_path = (
|
||||
payload.get("adapter_path")
|
||||
or payload.get("adapter_name_or_path")
|
||||
or (trained_model and (trained_model.get("artifact_dir") or trained_model.get("adapter_path") or trained_model.get("merged_path")))
|
||||
)
|
||||
if not base_model_path:
|
||||
raise fail(400, "base_model_path is required")
|
||||
if not adapter_path:
|
||||
raise fail(400, "adapter_path is required")
|
||||
node = store.schedule_node({**payload, "gpus": payload.get("gpus") or []})
|
||||
requested_node_id = payload.get("requested_node_id") or payload.get("compute_node_id") or (trained_model and trained_model.get("compute_node_id"))
|
||||
node = store.schedule_node({**payload, "requested_node_id": requested_node_id, "gpus": payload.get("gpus") or []})
|
||||
health = node.get("health_detail") or {}
|
||||
output_root = str(health.get("output_root") or f"{node['data_root'].rstrip('/')}/outputs")
|
||||
output_name = str(payload.get("output_model_name") or payload.get("merged_model_name") or f"{trained_model_id or 'model'}-merged")
|
||||
@@ -753,11 +804,13 @@ async def merge_model(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
"gpus": payload.get("gpus") or [],
|
||||
"trained_model_id": trained_model["id"] if trained_model else trained_model_id,
|
||||
"model_name": trained_model["name"] if trained_model else payload.get("model_name"),
|
||||
"compute_node_id": node["id"],
|
||||
"compute_node_code": node.get("code"),
|
||||
}
|
||||
if get_settings().compute_mode == "simulator":
|
||||
job = {"id": job_payload["id"], "status": "queued", "progress": 10, "command": [], "output_dir": output_dir}
|
||||
else:
|
||||
client = ComputeNodeClient(node["api_base_url"])
|
||||
client = ComputeNodeClient(node["api_base_url"], timeout=900)
|
||||
preview = await client.validate_job(job_payload)
|
||||
if not preview.get("valid", False):
|
||||
raise fail(409, "; ".join(preview.get("errors") or ["merge preflight failed"]))
|
||||
@@ -1537,28 +1590,49 @@ async def model_compare_detail(task_id: str) -> dict[str, Any]:
|
||||
raise fail(404, "compare task not found")
|
||||
|
||||
|
||||
async def _unload_from_compute_node() -> dict[str, Any]:
|
||||
"""Best-effort unload the inference model from the first online compute node."""
|
||||
store = get_platform_store()
|
||||
# Clear all inference tracking — only one model can be loaded at a time
|
||||
for node in store.compute_nodes():
|
||||
store.mark_inference_unloaded(node["id"])
|
||||
node = _select_first_online_node(store)
|
||||
if not node:
|
||||
return {"unloaded": False, "error": "no online compute node"}
|
||||
try:
|
||||
client = ComputeNodeClient(node["api_base_url"])
|
||||
result = await client._request("POST", "/inference/unload", json_data={})
|
||||
return result
|
||||
except Exception as exc:
|
||||
return {"unloaded": False, "error": str(exc)}
|
||||
async def _unload_from_compute_node(store: Any, task: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
"""Best-effort unload the inference model from the node(s) that hold it.
|
||||
|
||||
任务感知:优先卸载 ``task.load_status.loaded_models`` 中记录的节点;
|
||||
无任务时回退到平台记录的已加载推理的节点。每个节点使用短超时,
|
||||
保证卸载永远不会长时间阻塞调用方(例如删除操作)。
|
||||
"""
|
||||
node_ids: set[str] = set()
|
||||
if task:
|
||||
load_status = task.get("load_status") or {}
|
||||
if isinstance(load_status, str):
|
||||
try:
|
||||
load_status = json.loads(load_status)
|
||||
except json.JSONDecodeError:
|
||||
load_status = {}
|
||||
node_ids = {item.get("node_id") for item in load_status.get("loaded_models") or [] if item.get("node_id")}
|
||||
if not node_ids:
|
||||
node_ids = {node["id"] for node in store.compute_nodes() if store.is_inference_loaded(node["id"])}
|
||||
nodes = [node for node in store.compute_nodes() if node["id"] in node_ids]
|
||||
results: list[dict[str, Any]] = []
|
||||
for node in nodes:
|
||||
try:
|
||||
result = await ComputeNodeClient(node["api_base_url"]).inference_unload()
|
||||
results.append({"node_id": node["id"], "node_code": node.get("code"), "success": True, "result": result})
|
||||
except Exception as exc: # noqa: BLE001 - best-effort unload must not raise
|
||||
results.append({"node_id": node["id"], "node_code": node.get("code"), "success": False, "error": str(exc)})
|
||||
finally:
|
||||
store.mark_inference_unloaded(node["id"])
|
||||
return {"unloaded": bool(results), "nodes": results}
|
||||
|
||||
|
||||
@router.delete("/model-compare/{task_id}")
|
||||
async def model_compare_delete(task_id: str) -> dict[str, Any]:
|
||||
# 删除前先释放算力节点上的模型
|
||||
await _unload_from_compute_node()
|
||||
# 先删记录(快),再 best-effort 释放算力节点上的模型——删除绝不被卸载阻塞
|
||||
try:
|
||||
task = get_platform_store().compare_task(task_id)
|
||||
except KeyError:
|
||||
raise fail(404, "compare task not found")
|
||||
get_platform_store().delete_compare_task(task_id)
|
||||
try:
|
||||
await _unload_from_compute_node(get_platform_store(), task=task)
|
||||
except Exception: # noqa: BLE001 - deletion must succeed even if unload fails
|
||||
pass
|
||||
return ok({"deleted": task_id})
|
||||
|
||||
|
||||
@@ -1585,9 +1659,44 @@ async def model_compare_update_load_status(task_id: str, payload: dict[str, Any]
|
||||
raise fail(404, "compare task not found")
|
||||
|
||||
|
||||
def _invalidate_superseded_models(store: Any, task_id: str, loaded_models: list[dict[str, Any]]) -> None:
|
||||
"""同一计算节点同一时刻只能加载一个推理模型。
|
||||
|
||||
当新任务把模型派发到了某节点后,把其它任务中在该节点上 ready/running
|
||||
的模型标记为已被替换,保持平台 DB 与计算节点实际状态一致。
|
||||
"""
|
||||
taken_node_ids = {m.get("node_id") for m in loaded_models if m.get("node_id") and m.get("status") == "starting"}
|
||||
if not taken_node_ids:
|
||||
return
|
||||
for other in store.compare_tasks():
|
||||
if str(other.get("id")) == str(task_id):
|
||||
continue
|
||||
load_status = other.get("load_status") or {}
|
||||
if isinstance(load_status, str):
|
||||
try:
|
||||
load_status = json.loads(load_status)
|
||||
except json.JSONDecodeError:
|
||||
load_status = {}
|
||||
items = load_status.get("loaded_models") or []
|
||||
changed = False
|
||||
for item in items:
|
||||
if item.get("node_id") in taken_node_ids and item.get("status") in {"ready", "running"}:
|
||||
item["status"] = "error"
|
||||
item["error"] = "模型已被其他推理任务替换"
|
||||
changed = True
|
||||
if changed:
|
||||
new_status = "loaded" if any(i.get("status") in {"ready", "running"} for i in items) else "failed"
|
||||
store.update_compare_task(other["id"], {"status": new_status, "load_status": {"loaded_models": items}})
|
||||
|
||||
|
||||
@router.post("/model-compare/{task_id}/load")
|
||||
async def model_compare_load(task_id: str) -> dict[str, Any]:
|
||||
"""真正加载模型到算力节点(不再使用假 PID/端口)。"""
|
||||
"""异步派发模型加载到算力节点,立即返回。
|
||||
|
||||
加载进度由轮询对账器(compute_poller → reconcile_inference_loads)推进:
|
||||
任务项先以 status=starting 记录,对账器查询节点 /inference/status 后
|
||||
推进到 ready/error。这里只负责把加载请求派发出去,绝不同步等待加载完成。
|
||||
"""
|
||||
try:
|
||||
store = get_platform_store()
|
||||
task = store.compare_task(task_id)
|
||||
@@ -1597,15 +1706,14 @@ async def model_compare_load(task_id: str) -> dict[str, Any]:
|
||||
models = json.loads(models)
|
||||
except json.JSONDecodeError:
|
||||
models = []
|
||||
# 选取在线算力节点
|
||||
node = _select_first_online_node(store)
|
||||
if not node:
|
||||
online_nodes = _candidate_online_nodes(store)
|
||||
if not online_nodes:
|
||||
return ok({"status": "failed", "error": "no online compute node"})
|
||||
client = ComputeNodeClient(node["api_base_url"])
|
||||
loaded_models = []
|
||||
for item in models:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
preferred_node_id = item.get("node_id") or item.get("compute_node_id")
|
||||
model_path = item.get("model_path", "")
|
||||
if not model_path:
|
||||
# 尝试从模型库获取路径
|
||||
@@ -1614,28 +1722,46 @@ async def model_compare_load(task_id: str) -> dict[str, Any]:
|
||||
db_model = store.model(model_id)
|
||||
model_path = db_model.get("path", "")
|
||||
except KeyError:
|
||||
pass
|
||||
trained_model = next((m for m in store.trained_models() if str(m.get("id")) == str(model_id)), None)
|
||||
if trained_model:
|
||||
model_path = trained_model.get("merged_path") or trained_model.get("artifact_dir") or ""
|
||||
preferred_node_id = preferred_node_id or trained_model.get("compute_node_id")
|
||||
if not model_path:
|
||||
loaded_models.append({**item, "status": "error", "error": "model_path not found"})
|
||||
continue
|
||||
# 真正调用算力节点加载模型
|
||||
load_payload = {
|
||||
"model_name_or_path": model_path,
|
||||
"template": item.get("template", "qwen"),
|
||||
}
|
||||
if item.get("adapter_path"):
|
||||
load_payload["adapter_name_or_path"] = item["adapter_path"]
|
||||
try:
|
||||
result = await client._request("POST", "/inference/load", json_data=load_payload)
|
||||
if result.get("loaded"):
|
||||
if get_settings().compute_mode == "simulator":
|
||||
loaded_models.append({**item, "status": "ready", "node_id": "", "node_name": ""})
|
||||
continue
|
||||
# 只派发:HTTP 响应成功即视为已接受(节点会异步加载),loaded 字段忽略
|
||||
item_dispatched = False
|
||||
errors = []
|
||||
for node in _candidate_online_nodes(store, preferred_node_id):
|
||||
try:
|
||||
client = ComputeNodeClient(node["api_base_url"])
|
||||
await client.inference_load(load_payload)
|
||||
store.mark_inference_loaded(node["id"])
|
||||
loaded_models.append({**item, "status": "ready"})
|
||||
else:
|
||||
loaded_models.append({**item, "status": "error", "error": result.get("error", "load failed")})
|
||||
except Exception as exc:
|
||||
loaded_models.append({**item, "status": "error", "error": str(exc)})
|
||||
status = "loaded" if any(m.get("status") == "ready" for m in loaded_models) else "failed"
|
||||
loaded_models.append({**item, "status": "starting", "node_id": node["id"], "node_name": node.get("name")})
|
||||
item_dispatched = True
|
||||
break
|
||||
except Exception as exc: # noqa: BLE001 - try next candidate node
|
||||
errors.append(f"{node.get('name') or node.get('code')}: {exc}")
|
||||
if not item_dispatched:
|
||||
loaded_models.append({**item, "status": "error", "error": "; ".join(errors) or "load dispatch failed"})
|
||||
if any(m.get("status") == "starting" for m in loaded_models):
|
||||
status = "starting"
|
||||
elif any(m.get("status") == "error" for m in loaded_models):
|
||||
status = "failed"
|
||||
else:
|
||||
status = "loaded"
|
||||
updated = store.update_compare_task(task_id, {"status": status, "load_status": {"loaded_models": loaded_models}})
|
||||
# 同一节点同一时刻只能有一个推理模型;新任务占用了节点后,把其它任务上该节点的模型标记为已被替换
|
||||
_invalidate_superseded_models(store, task_id, loaded_models)
|
||||
return ok(updated)
|
||||
except KeyError:
|
||||
raise fail(404, "compare task not found")
|
||||
@@ -1645,8 +1771,9 @@ async def model_compare_load(task_id: str) -> dict[str, Any]:
|
||||
async def model_compare_unload(task_id: str) -> dict[str, Any]:
|
||||
try:
|
||||
store = get_platform_store()
|
||||
# 真正释放算力节点上的模型资源
|
||||
unload_result = await _unload_from_compute_node()
|
||||
task = store.compare_task(task_id)
|
||||
# 任务感知卸载:只释放该任务实际加载到的节点,短超时快速返回
|
||||
unload_result = await _unload_from_compute_node(store, task=task)
|
||||
updated = store.update_compare_task(task_id, {"status": "pending", "load_status": {"loaded_models": []}})
|
||||
return ok({"task": updated, "unload": unload_result})
|
||||
except KeyError:
|
||||
@@ -1662,7 +1789,7 @@ async def model_compare_start_model(task_id: str, payload: dict[str, Any] = Body
|
||||
async def model_compare_chat_with_port(payload: dict[str, Any] = Body(...)) -> dict[str, Any]:
|
||||
"""Proxy non-streaming chat to the compute node running the inference model."""
|
||||
store = get_platform_store()
|
||||
node = _select_first_online_node(store)
|
||||
node = _node_for_inference_payload(store, payload)
|
||||
if not node:
|
||||
return ok({"response": "no online compute node available for inference", "request": payload})
|
||||
try:
|
||||
@@ -1717,8 +1844,9 @@ async def model_chat_local_preload(payload: dict[str, Any] = Body(...)) -> dict[
|
||||
return ok({"loaded": False, "error": "no online compute node"})
|
||||
try:
|
||||
client = ComputeNodeClient(node["api_base_url"])
|
||||
result = await client._request("POST", "/inference/load", json_data=payload)
|
||||
if result.get("loaded"):
|
||||
# 计算节点现在异步加载:HTTP 接受(loading/ready)即视为派发成功
|
||||
result = await client.inference_load(payload)
|
||||
if result.get("loaded") or result.get("status") in {"loading", "ready"}:
|
||||
store.mark_inference_loaded(node["id"])
|
||||
return ok(result)
|
||||
except Exception as exc:
|
||||
@@ -1729,18 +1857,19 @@ async def model_chat_local_preload(payload: dict[str, Any] = Body(...)) -> dict[
|
||||
async def model_chat_local_unload() -> dict[str, Any]:
|
||||
"""Unload the inference model from the compute node."""
|
||||
store = get_platform_store()
|
||||
# Clear all inference tracking
|
||||
# 释放所有已加载推理的节点(短超时,best-effort)
|
||||
results: list[dict[str, Any]] = []
|
||||
for n in store.compute_nodes():
|
||||
store.mark_inference_unloaded(n["id"])
|
||||
node = _select_first_online_node(store)
|
||||
if not node:
|
||||
return ok({"unloaded": False, "error": "no online compute node"})
|
||||
try:
|
||||
client = ComputeNodeClient(node["api_base_url"])
|
||||
result = await client._request("POST", "/inference/unload", json_data={})
|
||||
return ok(result)
|
||||
except Exception as exc:
|
||||
return ok({"unloaded": False, "error": str(exc)})
|
||||
if not store.is_inference_loaded(n["id"]):
|
||||
continue
|
||||
try:
|
||||
result = await ComputeNodeClient(n["api_base_url"]).inference_unload()
|
||||
results.append({"node_id": n["id"], "success": True, "result": result})
|
||||
except Exception as exc: # noqa: BLE001 - best-effort unload
|
||||
results.append({"node_id": n["id"], "success": False, "error": str(exc)})
|
||||
finally:
|
||||
store.mark_inference_unloaded(n["id"])
|
||||
return ok({"unloaded": True, "nodes": results})
|
||||
|
||||
|
||||
@router.get("/model-chat/local/status")
|
||||
@@ -1752,7 +1881,7 @@ async def model_chat_local_status() -> dict[str, Any]:
|
||||
return ok({"loaded": False, "error": "no online compute node"})
|
||||
try:
|
||||
client = ComputeNodeClient(node["api_base_url"])
|
||||
result = await client._request("GET", "/inference/status")
|
||||
result = await client.inference_status()
|
||||
return ok(result)
|
||||
except Exception as exc:
|
||||
return ok({"loaded": False, "error": str(exc)})
|
||||
@@ -1770,8 +1899,9 @@ async def model_chat_trained_preload(payload: dict[str, Any] = Body(...)) -> dic
|
||||
return ok({"loaded": False, "error": "no online compute node"})
|
||||
try:
|
||||
client = ComputeNodeClient(node["api_base_url"])
|
||||
result = await client._request("POST", "/inference/load", json_data=payload)
|
||||
if result.get("loaded"):
|
||||
# 计算节点现在异步加载:HTTP 接受(loading/ready)即视为派发成功
|
||||
result = await client.inference_load(payload)
|
||||
if result.get("loaded") or result.get("status") in {"loading", "ready"}:
|
||||
store.mark_inference_loaded(node["id"])
|
||||
return ok(result)
|
||||
except Exception as exc:
|
||||
|
||||
Reference in New Issue
Block a user