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