daavidhauser's picture
Publish Swift HyperQwen collection with performance and quality comparisons
2bc6021 verified
Raw History Blame Contribute Delete
6.02 kB
"""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))