add minimal model compatibility preflight agent

This commit is contained in:
2026-08-25 21:13:19 +08:00
parent 4b28e819b4
commit b37946c30c
4 changed files with 356 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,51 @@
# xc-model-compat-agent
信创模盒模型兼容性预检智能体:评估模型规模、精度、显存与目标国产算力适配风险
信创模盒模型兼容性预检智能体。它根据用户提供的模型规模、精度、上下文长度、卡数和单卡显存,给出保守的显存预估、风险项与下一步验证建议
## 安全边界
- 只分析请求中的数据,不访问外部网络。
- 不读取或输出凭据。
- 不创建、提交或修改任何模型适配任务。
- 结论属于部署前预检,不替代目标芯片上的真实运行验证。
## 平台约束
- `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
```

247
main.py Normal file
View File

@@ -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()

48
test_main.py Normal file
View File

@@ -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()