"""Own one benchmark server process group; leave unrelated services alone.""" import contextlib import json import os import signal import subprocess import time import urllib.request from common import ROOT, RUN, write_json, sha256, stamp from evaluate import key @contextlib.contextmanager def server(model, tag, batch=False, int8=False): context_len = int(os.environ.get("SWIFT15_CONTEXT_LEN", "16384")) speculation = os.environ.get("SWIFT15_SPEC", "mtp") if speculation not in {"mtp", "off"}: raise ValueError("SWIFT15_SPEC must be mtp or off") write_json(RUN / "status.json", {"stage": tag + "-server", "state": "starting", "started": stamp()}) log = RUN / "logs" / (tag + "-server.log") log.parent.mkdir(parents=True, exist_ok=True) env = dict(os.environ, MODEL=str(model), HOST="127.0.0.1", PORT="18021", VISION="0", CTX="fast", SPEC=speculation, MAX_LEN=str(context_len), MAX_SEQS="8", GPU_UTIL="0.93", API_SERVERS="1", PREFIX_CACHE="0", INT8_ACT="int8" if int8 else "", INT8_LAYERS="mlp", OMP_NUM_THREADS="8", EXTRA_ARGS="--generation-config vllm", VLLM_OFFLOAD_KEEP_SHM="1") for name in ["VLLM_MARLIN_INPUT_DTYPE", "VLLM_MARLIN_INT8_INCLUDE_RE", "VLLM_PREFILL_ATTN"]: env.pop(name, None) if not int8 and (ROOT / ".env").exists(): for line in (ROOT / ".env").read_text().splitlines(): if line.startswith("INT8_ACT=") and line.split("=", 1)[1].strip(' "'): raise RuntimeError("Local .env enables INT8_ACT; remove that default before the controlled W4A16 comparison") launcher = ROOT / ("batch/start_qwen.sh" if batch else "single-user/start_qwen.sh") if batch: env["EXTRA_ARGS"] += " --attention-backend FLASH_ATTN --kv-cache-dtype bfloat16 --no-enable-prefix-caching" local_profile = os.environ.get("SWIFT15_SERVING_PROFILE") == "local-single-user" and not batch if local_profile: # Match the launcher's .env semantics, retaining all local runtime knobs. env = dict(os.environ) for line in (ROOT / ".env").read_text().splitlines(): if not line or line.startswith("#"): continue name, sep, value = line.removeprefix("export ").partition("=") if sep and name.isidentifier() and not env.get(name): env[name] = value.strip('"') context_len = 150000 env.update(MODEL=str(model), HOST="127.0.0.1", PORT="18021", CTX="long", MAX_LEN=str(context_len)) if env.get("SPEC", "mtp") != "mtp": raise ValueError("The FP8 local profile requires SPEC=mtp") speculation = "mtp" import socket with socket.socket() as s: if s.connect_ex(("127.0.0.1", 18021)) == 0: raise RuntimeError("Benchmark port 18021 is occupied; refusing to use or stop another server") with log.open("w") as out: p = subprocess.Popen(["bash", str(launcher)], cwd=ROOT, env=env, stdout=out, stderr=subprocess.STDOUT, start_new_session=True) identity = {"created": stamp(), "model": str(model), "config_sha256": sha256(model / "config.json"), "index_sha256": sha256(model / "model.safetensors.index.json"), "launcher_sha256": sha256(launcher), "tokenizer_sha256": sha256(model / "tokenizer.json"), "chat_template_sha256": sha256(model / "chat_template.jinja") if (model / "chat_template.jinja").exists() else sha256(model / "tokenizer_config.json"), "shards": {p.name: {"size": p.stat().st_size, "mtime_ns": p.stat().st_mtime_ns} for p in model.glob("*.safetensors")}, "mode": "batch" if batch else "single-user", "int8_activations": int8, "speculation": speculation, "max_model_len": context_len, "max_num_seqs": 8, "kv_cache_dtype": "bfloat16", "prefix_cache": False, "gpu_memory_utilization": .93, "pid": p.pid} if local_profile: identity.update(serving_profile="local-single-user", kv_cache_dtype="fp8", max_num_seqs=int(env.get("MAX_SEQS") or 8), prefix_cache=env.get("PREFIX_CACHE") == "1", vision=env.get("VISION") == "1", gpu_memory_utilization=float(env.get("GPU_UTIL") or .93), draft_tokens=int(env.get("DRAFT_TOKENS") or 3), int8_activations=bool(env.get("INT8_ACT")), local_env_sha256=sha256(ROOT / ".env"), extra_args=env.get("EXTRA_ARGS", "")) write_json(RUN / "active-server.json", identity) write_json(RUN / "results" / tag / "server.json", identity) try: deadline = time.monotonic() + 900 while time.monotonic() < deadline: if p.poll() is not None: raise RuntimeError(f"Server exited {p.returncode}; see {log}") try: req = urllib.request.Request("http://127.0.0.1:18021/health", headers={"Authorization": "Bearer " + key()}) with urllib.request.urlopen(req, timeout=3) as response: if response.status == 200: break except OSError: time.sleep(2) else: raise TimeoutError(f"Server startup timed out; see {log}") yield identity finally: try: os.killpg(p.pid, signal.SIGTERM) except ProcessLookupError: pass try: p.wait(timeout=45) except subprocess.TimeoutExpired: try: os.killpg(p.pid, signal.SIGKILL) except ProcessLookupError: pass p.wait(timeout=15) write_json(RUN / "active-server.json", dict(identity, stopped=stamp(), returncode=p.returncode))