add 寒武纪 MLU590 advisor agent
This commit is contained in:
12
Dockerfile
Normal file
12
Dockerfile
Normal 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"]
|
||||
56
README.md
56
README.md
@@ -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
234
main.py
Normal 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
50
test_main.py
Normal 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()
|
||||
Reference in New Issue
Block a user