Files

235 lines
8.4 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""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-biren-166m-advisor-agent"
AGENT_VERSION = "1.0.0"
CARD_VENDOR = "壁仞"
CARD_MODEL = "壁砺166M"
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()