Files

248 lines
8.8 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.

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