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 bc17afe..fea703b 100644 --- a/README.md +++ b/README.md @@ -1,3 +1,51 @@ # xc-model-compat-agent -信创模盒模型兼容性预检智能体:评估模型规模、精度、显存与目标国产算力适配风险。 \ No newline at end of file +信创模盒的模型兼容性预检智能体。它根据用户提供的模型规模、精度、上下文长度、卡数和单卡显存,给出保守的显存预估、风险项与下一步验证建议。 + +## 安全边界 + +- 只分析请求中的数据,不访问外部网络。 +- 不读取或输出凭据。 +- 不创建、提交或修改任何模型适配任务。 +- 结论属于部署前预检,不替代目标芯片上的真实运行验证。 + +## 平台约束 + +- `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/7b-model", + "architecture": "transformer", + "parameters_b": 7, + "dtype": "bf16", + "hardware": "目标国产算力卡", + "device_memory_gb": 24, + "cards": 1, + "context_length": 8192 + }' +``` + +## 本地验证 + +```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..bd8e18a --- /dev/null +++ b/main.py @@ -0,0 +1,247 @@ +"""ModelHub XC model compatibility preflight agent. + +The service is intentionally read-only: it evaluates supplied model and device +facts but never submits adaptation jobs or calls external services. +""" + +from __future__ import annotations + +import json +import os +import signal +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Any + + +AGENT_NAME = "xc-model-compat-agent" +AGENT_VERSION = "1.0.0" +PORT = int(os.getenv("PORT", "8080")) +STRATEGY_ID = os.getenv("STRATEGY_ID", "") +MAX_BODY_BYTES = 1_000_000 + +DTYPE_BYTES = { + "fp32": 4.0, + "float32": 4.0, + "fp16": 2.0, + "float16": 2.0, + "bf16": 2.0, + "bfloat16": 2.0, + "int8": 1.0, + "8bit": 1.0, + "int4": 0.5, + "4bit": 0.5, +} + + +class ValidationError(ValueError): + """Raised when a provided 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_float(value: Any, field: str) -> float | None: + if value in (None, ""): + return None + try: + number = float(value) + except (TypeError, ValueError) as exc: + raise ValidationError(f"{field} 必须是数字") from exc + if number <= 0: + raise ValidationError(f"{field} 必须大于 0") + return number + + +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 analyze(payload: dict[str, Any]) -> dict[str, Any]: + """Return a conservative, evidence-labelled compatibility preflight.""" + + if not isinstance(payload, dict): + raise ValidationError("请求正文必须是 JSON 对象") + + model_name = str( + _first(payload, "model_name", "model", "model_address") or "未指定模型" + ) + architecture = str(_first(payload, "architecture", "model_architecture") or "") + hardware = str(_first(payload, "hardware", "gpu_type", "target_gpu") or "") + dtype = str(_first(payload, "dtype", "precision") or "bf16").lower() + parameters_b = _positive_float( + _first(payload, "parameters_b", "parameter_billions"), "parameters_b" + ) + device_memory_gb = _positive_float( + _first(payload, "device_memory_gb", "memory_gb"), "device_memory_gb" + ) + cards = _positive_int(_first(payload, "cards", "card_count"), "cards", 1) + context_length = _positive_int( + _first(payload, "context_length", "max_sequence_length"), + "context_length", + 4096, + ) + + missing_fields: list[str] = [] + if parameters_b is None: + missing_fields.append("parameters_b") + if device_memory_gb is None: + missing_fields.append("device_memory_gb") + if not hardware: + missing_fields.append("hardware") + + risks: list[str] = [] + recommendations: list[str] = [] + estimates: dict[str, float] = {} + verdict = "information_required" + + bytes_per_parameter = DTYPE_BYTES.get(dtype) + if bytes_per_parameter is None: + risks.append(f"精度 {dtype} 未进入估算表,需确认格式与推理后端支持。") + elif dtype in {"int8", "8bit", "int4", "4bit"}: + risks.append("量化精度需要目标国产芯片后端提供对应量化算子。") + + if not architecture: + risks.append("模型架构未提供,无法预判注意力、归一化及自定义算子兼容性。") + if context_length > 32768: + risks.append("长上下文会显著增加 KV Cache;当前权重估算未包含业务并发负载。") + if cards > 1: + risks.append("多卡方案需额外验证张量并行、集合通信和卡间带宽。") + + if parameters_b is not None and bytes_per_parameter is not None: + weights_gib = parameters_b * 1_000_000_000 * bytes_per_parameter / (1024**3) + runtime_budget_gib = weights_gib * 1.20 + estimates = { + "weights_gib": round(weights_gib, 2), + "minimum_runtime_budget_gib": round(runtime_budget_gib, 2), + } + if device_memory_gb is not None: + total_device_memory_gb = device_memory_gb * cards + estimates["total_device_memory_gb"] = round(total_device_memory_gb, 2) + estimates["headroom_gb"] = round(total_device_memory_gb - runtime_budget_gib, 2) + if total_device_memory_gb < runtime_budget_gib: + verdict = "insufficient_memory" + recommendations.append("降低精度、增加卡数或选择更大显存设备后再做算子验证。") + elif runtime_budget_gib > total_device_memory_gb * 0.85: + verdict = "tight_memory" + recommendations.append("显存余量偏小,先用低并发和较短上下文完成冒烟测试。") + else: + verdict = "preflight_feasible" + recommendations.append("显存预检通过,可进入目标后端算子兼容性测试。") + + recommendations.extend( + [ + "确认目标芯片 SDK、PyTorch 后端与推理框架版本组合。", + "检查 Attention、RMSNorm/LayerNorm、RoPE 和自定义算子支持情况。", + "先跑单请求冒烟测试,再分别测首 Token 延迟、吞吐和峰值显存。", + ] + ) + + return { + "agent": AGENT_NAME, + "version": AGENT_VERSION, + "model": model_name, + "target": { + "hardware": hardware or None, + "cards": cards, + "device_memory_gb": device_memory_gb, + }, + "inputs": { + "architecture": architecture or None, + "parameters_b": parameters_b, + "dtype": dtype, + "context_length": context_length, + }, + "verdict": verdict, + "missing_fields": missing_fields, + "estimates": estimates, + "risks": risks, + "recommendations": recommendations, + "disclaimer": "这是部署前预检,不代表目标芯片已完成真实运行验证。", + } + + +class Handler(BaseHTTPRequestHandler): + server_version = "ModelHubCompatAgent/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": "信创模型兼容性与显存预检", + "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..dec2722 --- /dev/null +++ b/test_main.py @@ -0,0 +1,48 @@ +import unittest + +from main import ValidationError, analyze + + +class AnalyzeTests(unittest.TestCase): + def test_feasible_preflight(self): + result = analyze( + { + "model_name": "example/7b-model", + "architecture": "transformer", + "parameters_b": 7, + "dtype": "bf16", + "hardware": "国产算力卡", + "device_memory_gb": 24, + "cards": 1, + } + ) + self.assertEqual(result["verdict"], "preflight_feasible") + self.assertGreater(result["estimates"]["headroom_gb"], 0) + + def test_insufficient_memory(self): + result = analyze( + { + "model_name": "example/32b-model", + "architecture": "transformer", + "parameters_b": 32, + "dtype": "bf16", + "hardware": "国产算力卡", + "device_memory_gb": 24, + "cards": 2, + } + ) + self.assertEqual(result["verdict"], "insufficient_memory") + + def test_missing_information_is_explicit(self): + result = analyze({"model_name": "example/unknown-model"}) + self.assertEqual(result["verdict"], "information_required") + self.assertIn("parameters_b", result["missing_fields"]) + self.assertIn("device_memory_gb", result["missing_fields"]) + + def test_invalid_card_count_is_rejected(self): + with self.assertRaises(ValidationError): + analyze({"cards": 0}) + + +if __name__ == "__main__": + unittest.main()