add 寒武纪 MLU590 advisor agent

This commit is contained in:
2026-08-25 22:58:53 +08:00
parent d8c3dac8db
commit 2e86a9e78e
4 changed files with 351 additions and 1 deletions

12
Dockerfile Normal file
View File

@@ -0,0 +1,12 @@
FROM modelhubxc-4pd.tencentcloudcr.com/xc_agent_platform/python:3.11-slim
ENV PYTHONDONTWRITEBYTECODE=1 \
PYTHONUNBUFFERED=1 \
PORT=8080
WORKDIR /app
COPY main.py /app/main.py
EXPOSE 8080
CMD ["python", "/app/main.py"]

View File

@@ -1,3 +1,57 @@
# xc-mlu590-advisor-agent
寒武纪 MLU590 信创模盒专属部署前预检智能体;校验目标、版本矩阵和运行风险
信创模盒的 **寒武纪MLU590** 专属部署前预检智能体。
## 能做什么
- 校验用户指定的目标卡是否与本智能体的官方卡型一致。
- 收集模型、框架、后端、驱动和 SDK 版本,明确缺失信息。
- 根据多卡、量化、长上下文、自定义算子和动态形状提示验证风险。
- 给出从设备可见性到模型最小输入的验证顺序。
## 事实边界
- 卡型名称来自信创模盒“模型 X 算力”页面。
- 不声明页面未提供的显存、算力、SDK 版本或算子支持结论。
- 不访问外部网络,不执行命令,不创建或修改适配任务。
- 最终兼容性必须以目标设备上的真实日志和测试结果为准。
## 平台约束
- 仓库根目录包含 `Dockerfile`,容器开放 `8080`
- `GET /health` 返回 HTTP 200。
- 读取 `STRATEGY_ID`,仅报告是否存在,不输出其值。
- 处理 `SIGTERM`,在平台 30 秒窗口内退出。
- 仅使用 Python 标准库。
## 接口
- `GET /health`:存活检查。
- `GET /`:卡型和智能体元信息。
- `POST /analyze`:部署前预检。
- `POST /task`:与 `/analyze` 相同的兼容入口。
示例:
```bash
curl -X POST http://localhost:8080/analyze \
-H 'Content-Type: application/json' \
-d '{
"model_name": "example/model",
"hardware": "寒武纪 MLU590",
"framework": "PyTorch",
"backend": "目标后端",
"sdk_version": "实际版本",
"driver_version": "实际版本",
"precision": "bf16",
"cards": 1
}'
```
## 本地验证
```bash
python3 -m unittest -v
python3 main.py
curl http://localhost:8080/health
```

234
main.py Normal file
View File

@@ -0,0 +1,234 @@
"""Card-specific ModelHub XC deployment preflight agent.
This service is read-only and only makes claims supported by submitted facts.
"""
from __future__ import annotations
import json
import os
import re
import signal
import threading
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import Any
AGENT_NAME = "xc-mlu590-advisor-agent"
AGENT_VERSION = "1.0.0"
CARD_VENDOR = "寒武纪"
CARD_MODEL = "MLU590"
PORT = int(os.getenv("PORT", "8080"))
STRATEGY_ID = os.getenv("STRATEGY_ID", "")
MAX_BODY_BYTES = 1_000_000
class ValidationError(ValueError):
"""Raised when a request field is invalid."""
def _first(payload: dict[str, Any], *names: str) -> Any:
for name in names:
if name in payload and payload[name] not in (None, ""):
return payload[name]
return None
def _positive_int(value: Any, field: str, default: int) -> int:
if value in (None, ""):
return default
try:
number = int(value)
except (TypeError, ValueError) as exc:
raise ValidationError(f"{field} 必须是整数") from exc
if number <= 0:
raise ValidationError(f"{field} 必须大于 0")
return number
def _boolean(value: Any, field: str, default: bool = False) -> bool:
if value in (None, ""):
return default
if isinstance(value, bool):
return value
if isinstance(value, str):
normalized = value.strip().lower()
if normalized in {"true", "1", "yes", "on"}:
return True
if normalized in {"false", "0", "no", "off"}:
return False
raise ValidationError(f"{field} 必须是布尔值")
def _normalize(value: str) -> str:
return re.sub(r"[^0-9a-z一-鿿]+", "", value.lower())
def analyze(payload: dict[str, Any]) -> dict[str, Any]:
"""Return an evidence-labelled preflight for this repository's card."""
if not isinstance(payload, dict):
raise ValidationError("请求正文必须是 JSON 对象")
model_name = str(_first(payload, "model_name", "model", "model_address") or "")
framework = str(_first(payload, "framework") or "")
backend = str(_first(payload, "backend", "inference_engine", "engine") or "")
sdk_version = str(_first(payload, "sdk_version", "sdk") or "")
driver_version = str(_first(payload, "driver_version", "driver") or "")
precision = str(_first(payload, "precision", "dtype") or "bf16").lower()
cards = _positive_int(_first(payload, "cards", "card_count"), "cards", 1)
context_length = _positive_int(
_first(payload, "context_length", "max_sequence_length"),
"context_length",
4096,
)
custom_ops = _boolean(_first(payload, "custom_ops"), "custom_ops")
dynamic_shapes = _boolean(
_first(payload, "dynamic_shapes"), "dynamic_shapes"
)
supplied_hardware = str(
_first(payload, "hardware", "device", "target_card", "target_gpu") or ""
)
target_tokens = {_normalize(CARD_VENDOR), _normalize(CARD_MODEL)}
normalized_hardware = _normalize(supplied_hardware)
target_matches = not supplied_hardware or any(
token and token in normalized_hardware for token in target_tokens
)
environment = {
"model_name": model_name,
"framework": framework,
"backend": backend,
"sdk_version": sdk_version,
"driver_version": driver_version,
}
missing_fields = [name for name, value in environment.items() if not value]
risks: list[str] = []
if precision in {"int4", "4bit", "int8", "8bit"}:
risks.append("量化方案需要实测目标卡后端是否提供对应权重格式与算子内核。")
if cards > 1:
risks.append("多卡运行需要验证集合通信、进程数、拓扑和并行策略。")
if context_length > 32768:
risks.append("长上下文会增加 KV Cache 压力,需要按真实并发测峰值内存。")
if custom_ops:
risks.append("模型包含自定义算子需要确认编译链、ABI 与目标后端注册情况。")
if dynamic_shapes:
risks.append("动态形状需要验证图编译缓存、回退路径与重复编译开销。")
if not target_matches:
verdict = "target_mismatch"
risks.insert(0, f"本智能体仅面向 {CARD_VENDOR}{CARD_MODEL}")
elif missing_fields:
verdict = "information_required"
else:
verdict = "preflight_ready"
recommendations = [
f"确认实际目标设备标识为 {CARD_VENDOR}{CARD_MODEL}",
"记录驱动、SDK、框架和推理后端的完整版本矩阵。",
"先用厂商基础样例确认设备可见,再运行模型最小输入。",
"依次验证模型加载、首个算子、单请求输出和资源峰值。",
"真实支持性结论必须来自目标设备运行日志与健康检查。",
]
return {
"agent": AGENT_NAME,
"version": AGENT_VERSION,
"official_target": {
"vendor": CARD_VENDOR,
"model": CARD_MODEL,
"source_scope": "信创模盒模型 X 算力页面",
},
"supplied_hardware": supplied_hardware or None,
"target_matches": target_matches,
"verdict": verdict,
"inputs": {
**{key: value or None for key, value in environment.items()},
"precision": precision,
"cards": cards,
"context_length": context_length,
"custom_ops": custom_ops,
"dynamic_shapes": dynamic_shapes,
},
"missing_fields": missing_fields,
"risks": risks,
"recommendations": recommendations,
"disclaimer": "未提供或未实测的硬件规格与兼容性不会被推断为已支持。",
}
class Handler(BaseHTTPRequestHandler):
server_version = "ModelHubCardAdvisor/1.0"
def _json(self, payload: dict[str, Any], status: int = 200) -> None:
body = json.dumps(payload, ensure_ascii=False).encode("utf-8")
self.send_response(status)
self.send_header("Content-Type", "application/json; charset=utf-8")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def do_GET(self) -> None: # noqa: N802 - BaseHTTPRequestHandler contract
if self.path == "/health":
self._json({"status": "ok", "agent": AGENT_NAME, "version": AGENT_VERSION})
return
if self.path == "/":
self._json(
{
"name": AGENT_NAME,
"version": AGENT_VERSION,
"description": f"{CARD_VENDOR} {CARD_MODEL} 专属部署前预检",
"official_target": {"vendor": CARD_VENDOR, "model": CARD_MODEL},
"strategy_id_present": bool(STRATEGY_ID),
"endpoints": ["GET /health", "POST /analyze", "POST /task"],
"external_writes": False,
}
)
return
self._json({"error": "not_found"}, 404)
def do_POST(self) -> None: # noqa: N802 - BaseHTTPRequestHandler contract
if self.path not in {"/analyze", "/task"}:
self._json({"error": "not_found"}, 404)
return
try:
length = int(self.headers.get("Content-Length", "0"))
if length <= 0:
raise ValidationError("请求正文不能为空")
if length > MAX_BODY_BYTES:
self._json({"error": "payload_too_large"}, 413)
return
payload = json.loads(self.rfile.read(length))
self._json(analyze(payload))
except (json.JSONDecodeError, UnicodeDecodeError):
self._json({"error": "invalid_json", "message": "请求正文必须是有效 JSON"}, 400)
except ValidationError as exc:
self._json({"error": "validation_error", "message": str(exc)}, 422)
def log_message(self, _format: str, *_args: Any) -> None:
return
STOP = threading.Event()
def _handle_signal(signum: int, _frame: Any) -> None:
print(f"received signal {signum}; shutting down", flush=True)
STOP.set()
def main() -> None:
signal.signal(signal.SIGTERM, _handle_signal)
signal.signal(signal.SIGINT, _handle_signal)
server = ThreadingHTTPServer(("0.0.0.0", PORT), Handler)
server.timeout = 1
print(f"{AGENT_NAME} {AGENT_VERSION} listening on 0.0.0.0:{PORT}", flush=True)
while not STOP.is_set():
server.handle_request()
server.server_close()
if __name__ == "__main__":
main()

50
test_main.py Normal file
View File

@@ -0,0 +1,50 @@
import unittest
from main import CARD_MODEL, CARD_VENDOR, ValidationError, analyze
class AnalyzeTests(unittest.TestCase):
def complete_payload(self):
return {
"model_name": "example/model",
"hardware": f"{CARD_VENDOR} {CARD_MODEL}",
"framework": "PyTorch",
"backend": "vendor-backend",
"sdk_version": "provided-by-user",
"driver_version": "provided-by-user",
}
def test_complete_preflight_targets_this_card(self):
result = analyze(self.complete_payload())
self.assertEqual(result["verdict"], "preflight_ready")
self.assertTrue(result["target_matches"])
self.assertEqual(result["official_target"]["model"], CARD_MODEL)
def test_mismatched_target_is_rejected(self):
payload = self.complete_payload()
payload["hardware"] = "其他厂商 其他卡型"
result = analyze(payload)
self.assertEqual(result["verdict"], "target_mismatch")
self.assertFalse(result["target_matches"])
def test_missing_information_is_explicit(self):
result = analyze({})
self.assertEqual(result["verdict"], "information_required")
self.assertIn("sdk_version", result["missing_fields"])
self.assertIn("driver_version", result["missing_fields"])
def test_multicard_and_quantization_risks_are_reported(self):
payload = self.complete_payload()
payload.update({"cards": 2, "precision": "int8"})
result = analyze(payload)
joined = " ".join(result["risks"])
self.assertIn("多卡", joined)
self.assertIn("量化", joined)
def test_invalid_card_count_is_rejected(self):
with self.assertRaises(ValidationError):
analyze({"cards": 0})
if __name__ == "__main__":
unittest.main()