diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..3cf955d --- /dev/null +++ b/Dockerfile @@ -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"] diff --git a/README.md b/README.md index a0a1b8f..b9823cd 100644 --- a/README.md +++ b/README.md @@ -1,3 +1,57 @@ # xc-mlu590-advisor-agent -寒武纪 MLU590 信创模盒专属部署前预检智能体;校验目标、版本矩阵和运行风险。 \ No newline at end of file +信创模盒的 **寒武纪|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 +``` diff --git a/main.py b/main.py new file mode 100644 index 0000000..38db8fb --- /dev/null +++ b/main.py @@ -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() diff --git a/test_main.py b/test_main.py new file mode 100644 index 0000000..ce0c289 --- /dev/null +++ b/test_main.py @@ -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()