File size: 6,020 Bytes
2bc6021
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
"""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))