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