Files

248 lines
8.8 KiB
Python
Raw Permalink Normal View History

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