File size: 6,620 Bytes
964cd4b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
"""Decision Index engine for KnowLine (PelaAI): the model repo's `knowline_server.py` front end, in process.

One command, no server to start by hand:

    # default: the setting of our submitted runs. Starts SGLang with the flags of serve_knowline.sh (FP8 at load) on a
    # free local port, scores through it, and stops it when the run ends.
    python -m decision_index pipeline --engine knowline_engine:KnowLine \\
        --option model=PelaAI/KnowLine-4B-Gen2 --option revision=<commit> --out runs/KnowLine-4B-Gen2

    # without SGLang: transformers only, bf16 (slower; not the setting of our runs)
    ... --option backend=hf

Put this file and `knowline_server.py` (both in the model repo) on PYTHONPATH, e.g. run from a clone of the model repo.
Requirements: `transformers` and `requests`; `sglang==0.5.21` for the default backend, `torch` for backend=hf.
Options: model, revision, backend (sglang | hf), gpu (CUDA device for SGLang; default: as CUDA_VISIBLE_DEVICES), mem (0.72), port (free one),
temperature (1.0), workers (16), startup_timeout (s, 1800), sglang_python (interpreter with SGLang installed, if it is
not the one running the kit), device (backend=hf: torch device or device map, default auto / cuda).
Licence of this file: MIT.
"""

import atexit
import os
import signal
import socket
import subprocess
import sys
import time
from pathlib import Path

from decision_index.engines.base import Engine, Unsupported

sys.path.insert(0, str(Path(__file__).resolve().parent))
import knowline_server as ks  # noqa: E402

SGLANG_FLAGS = ["--served-model-name", "m", "--tp", "1", "--quantization", "fp8", "--mamba-radix-cache-strategy",
                "extra_buffer", "--enable-fp32-lm-head"]


def _free_port():
    with socket.socket() as s:
        s.bind(("127.0.0.1", 0))
        return s.getsockname()[1]


class KnowLine(Engine):
    name = "knowline"
    latency = ("In-process request wall time through knowline_server's KnowLine engine (rendering + one prefill per "
               "question), against a local SGLang server for backend=sglang; excludes model loading and server startup.")

    def __init__(self, model, revision=None, backend="sglang", gpu=None, mem=0.72, port=None, temperature=1.0,
                 workers=16, startup_timeout=1800, sglang_python=None, device=None, **options):
        super().__init__(**options)
        import requests
        import transformers
        from transformers import AutoTokenizer

        path = model
        if not Path(model).exists():  # a Hub repo id: pin the files once so SGLang and the tokenizer read the same ones
            from huggingface_hub import snapshot_download
            path = snapshot_download(model, revision=revision)
        self.model_id, self.backend_name, self.server = model, backend, None
        tok = AutoTokenizer.from_pretrained(path)
        if backend == "sglang":
            port = int(port or _free_port())
            env = dict(os.environ)
            if gpu is not None:
                env["CUDA_VISIBLE_DEVICES"] = str(gpu)
            cmd = [sglang_python or sys.executable, "-m", "sglang.launch_server", "--model-path", path, *SGLANG_FLAGS,
                   "--mem-fraction-static", str(mem), "--port", str(port)]
            self.server = subprocess.Popen(cmd, env=env, start_new_session=True)
            atexit.register(self.close)
            url, deadline = f"http://127.0.0.1:{port}", time.time() + float(startup_timeout)
            while True:
                if self.server.poll() is not None:
                    raise RuntimeError(f"SGLang exited with code {self.server.returncode} before becoming healthy")
                try:
                    if requests.get(f"{url}/health", timeout=3).status_code == 200:
                        break
                except requests.RequestException:
                    pass
                if time.time() > deadline:
                    raise RuntimeError("SGLang did not become healthy in time")
                time.sleep(5)
            back = ks.SGLang(url)
        elif backend == "hf":
            back = ks.HF(path, device=device)
        else:
            raise ValueError("backend must be 'sglang' or 'hf'")
        self.engine = ks.KnowLine(tok, back, float(temperature), int(workers), {})
        self.provenance = {
            "kind": f"knowline_server.KnowLine in process, backend {backend}",
            "repo": model, "revision": revision, "front_end": "knowline_server.py (chat style, label-token softmax, "
            f"temperature {temperature}, {workers} scoring threads, no calibration file)",
            "sglang": " ".join(SGLANG_FLAGS + ["--mem-fraction-static", str(mem)]) if backend == "sglang" else None,
            "transformers": transformers.__version__,
            "policy": f"Unmodified state and questions; up to {ks.MAX_QUESTIONS} questions per request and "
                      f"{ks.MAX_LABELS} options per question, larger requests are unsupported, nothing is truncated.",
        }

    def __call__(self, state, questions):
        if len(questions) > ks.MAX_QUESTIONS:
            raise Unsupported(f"{len(questions)} questions; at most {ks.MAX_QUESTIONS} per request")
        try:
            answers, usage = self.engine.run(state, questions)
        except ValueError as exc:
            if any(k in str(exc) for k in ("criteria", "levels", "options", "questions")):
                raise Unsupported(str(exc)) from exc
            raise
        return {"model": self.model_id, "answers": answers, "usage": usage}, None

    def runtime(self):
        info = {"backend": self.backend_name}
        try:
            import torch
            info.update(torch=torch.__version__, cuda=torch.version.cuda)
            if torch.cuda.is_available():
                info["gpu"] = torch.cuda.get_device_name()
        except ImportError:
            pass
        if self.backend_name == "sglang":
            try:
                import sglang
                info["sglang"] = sglang.__version__
            except Exception:  # noqa: BLE001 - version is informational
                pass
        return info

    def close(self):
        if self.server is not None and self.server.poll() is None:
            try:
                os.killpg(self.server.pid, signal.SIGTERM)
                self.server.wait(timeout=60)
            except Exception:  # noqa: BLE001 - make sure the server does not outlive the run
                try:
                    os.killpg(self.server.pid, signal.SIGKILL)
                except ProcessLookupError:
                    pass
        self.server = None