235 lines
8.4 KiB
Python
235 lines
8.4 KiB
Python
"""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-ascend-950pr-advisor-agent"
|
||
AGENT_VERSION = "1.0.0"
|
||
CARD_VENDOR = "昇腾"
|
||
CARD_MODEL = "950PR"
|
||
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()
|