248 lines
8.8 KiB
Python
248 lines
8.8 KiB
Python
"""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()
|