add minimal model compatibility preflight 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"]
|
||||
50
README.md
50
README.md
@@ -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
247
main.py
Normal 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
48
test_main.py
Normal 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()
|
||||
Reference in New Issue
Block a user