File size: 8,565 Bytes
faea13d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1e54449
 
 
 
 
 
 
 
 
 
ff122a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
faea13d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
#!/usr/bin/env python3
"""對已經在跑的 llama-server 打一次 chat request,同時量記憶體。

量到的东西(每 0.2s 取樣一次,取峰值):
  - total RSS      /proc/<pid>/status 的 VmRSS(**包含 file-backed mmap 的權重頁**)
  - anon RSS       smaps_rollup 的 Anonymous,扣掉 page cache 後的「真的在用 RAM」
  - file RSS       smaps_rollup 的 Private_Dirty + Private_Clean(權重對映)
  - swap           /proc/<pid>/status 的 VmSwap(必須是 0)
  - io             /proc/<pid>/io 的 read_bytes(真正落到 block device 的量)

為什麼只看 total RSS:GGUF 是 mmap(MAP_SHARED) 對映的,權重頁會算進 RSS
而且真的佔用系統 RAM。只看匿名記憶體會嚴重低估。這個坑本專案第一版就踩過。

用法:
    python3 tools/verify_run.py --port 8080 --pid 12345 --model path.gguf --out out.json
"""
import argparse
import json
import os
import subprocess
import sys
import threading
import time
import urllib.error
import urllib.request

MIB = 1024 * 1024


def read_status(pid: int) -> dict:
    out = {}
    try:
        with open(f"/proc/{pid}/status") as f:
            for line in f:
                k, _, v = line.partition(":")
                if k in ("VmRSS", "VmSwap", "VmHWM"):
                    out[k] = int(v.split()[0]) * 1024
    except (OSError, ValueError, IndexError):
        pass
    return out


def read_smaps(pid: int) -> dict:
    out = {"Anonymous": 0, "Private_Dirty": 0, "Private_Clean": 0, "Shared_Clean": 0}
    try:
        with open(f"/proc/{pid}/smaps_rollup") as f:
            for line in f:
                k, _, v = line.partition(":")
                if k in out:
                    out[k] = int(v.split()[0]) * 1024
    except (OSError, ValueError, IndexError):
        pass
    return out


def read_io(pid: int) -> dict:
    out = {}
    try:
        with open(f"/proc/{pid}/io") as f:
            for line in f:
                k, _, v = line.partition(":")
                out[k] = int(v.strip())
    except (OSError, ValueError):
        pass
    return out


class Sampler(threading.Thread):
    def __init__(self, pid: int, interval: float = 0.2):
        super().__init__(daemon=True)
        self.pid = pid
        self.interval = interval
        self.stop_flag = threading.Event()
        self.peak = {"total_rss": 0, "anon_rss": 0, "file_rss": 0, "swap": 0, "hwm_rss": 0}
        self.io0 = None
        self.io1 = None

    def run(self):
        self.io0 = read_io(self.pid)
        while not self.stop_flag.is_set():
            st = read_status(self.pid)
            sm = read_smaps(self.pid)
            anon = sm["Anonymous"]
            filed = sm["Private_Dirty"] + sm["Private_Clean"]
            self.peak["total_rss"] = max(self.peak["total_rss"], st.get("VmRSS", 0))
            self.peak["anon_rss"] = max(self.peak["anon_rss"], anon)
            self.peak["file_rss"] = max(self.peak["file_rss"], filed)
            self.peak["swap"] = max(self.peak["swap"], st.get("VmSwap", 0))
            self.peak["hwm_rss"] = max(self.peak["hwm_rss"], st.get("VmHWM", 0))
            self.stop_flag.wait(self.interval)
        self.io1 = read_io(self.pid)


def wait_health(port: int, timeout: float) -> bool:
    end = time.time() + timeout
    while time.time() < end:
        try:
            with urllib.request.urlopen(f"http://127.0.0.1:{port}/health", timeout=5) as r:
                if r.status == 200:
                    return True
        except (urllib.error.URLError, OSError):
            time.sleep(2)
    return False


def chat(port: int, prompt: str, max_tokens: int, timeout: float) -> dict:
    body = json.dumps({
        "messages": [{"role": "user", "content": prompt}],
        "max_tokens": max_tokens,
        "temperature": 0.0,
    }).encode()
    req = urllib.request.Request(
        f"http://127.0.0.1:{port}/v1/chat/completions",
        data=body, headers={"Content-Type": "application/json"}, method="POST")
    t0 = time.time()
    with urllib.request.urlopen(req, timeout=timeout) as r:
        data = json.loads(r.read())
    dt = time.time() - t0
    usage = data.get("usage", {})
    n = usage.get("completion_tokens", 0)
    return {
        "prompt_tokens": usage.get("prompt_tokens"),
        "completion_tokens": n,
        "seconds": round(dt, 2),
        "tok_per_s": round(n / dt, 4) if dt > 0 and n else None,
        "content": data["choices"][0]["message"]["content"],
        "timings": data.get("timings"),
    }


def main() -> int:
    ap = argparse.ArgumentParser()
    ap.add_argument("--port", type=int, required=True)
    ap.add_argument("--pid", type=int, required=True)
    ap.add_argument("--model", required=True)
    ap.add_argument("--out", required=True)
    ap.add_argument("--prompt", default="Explain in one sentence what a MoE layer does.")
    ap.add_argument("--max-tokens", type=int, default=16)
    ap.add_argument("--timeout", type=float, default=1800)
    ap.add_argument("--ram-budget-mb", type=int, default=0)
    ap.add_argument("--arena-mb", type=int, default=0)
    ap.add_argument("--expert-slots", type=int, default=0)
    a = ap.parse_args()

    result = {
        "when": time.strftime("%Y-%m-%dT%H:%M:%S%z"),
        "model": a.model,
        "port": a.port,
        "pid": a.pid,
        "ram_budget_mb": a.ram_budget_mb,
        "arena_mb": a.arena_mb,
        "expert_slots": a.expert_slots,
    }

    if not wait_health(a.port, 300):
        result["error"] = "server 沒有在 300 秒內 ready"
        _write(a.out, result)
        return 1

    s = Sampler(a.pid)
    s.start()
    try:
        result["chat"] = chat(a.port, a.prompt, a.max_tokens, a.timeout)
    except (urllib.error.URLError, OSError, ValueError, KeyError) as e:
        result["error"] = f"chat failed: {e}"
    time.sleep(2)  # 讓峰值取樣涵蓋到收尾
    s.stop_flag.set()
    s.join(timeout=10)

    io_delta = {}
    if s.io0 and s.io1:
        for k in ("read_bytes", "rchar", "write_bytes"):
            if k in s.io0 and k in s.io1:
                io_delta[k] = s.io1[k] - s.io0[k]

    result["peak"] = {
        "total_rss_gb": round(s.peak["total_rss"] / 1024 ** 3, 4),
        "anon_rss_gb": round(s.peak["anon_rss"] / 1024 ** 3, 4),
        "file_rss_gb": round(s.peak["file_rss"] / 1024 ** 3, 4),
        "peak_swap_gb": round(s.peak["swap"] / 1024 ** 3, 4),
        "hwm_rss_gb": round(s.peak["hwm_rss"] / 1024 ** 3, 4),
    }
    result["io_delta"] = io_delta

    # 分頁統計(patches/0001 的 ST_STATS_FILE)
    stats_path = os.environ.get("ST_STATS_FILE")
    if stats_path and os.path.exists(stats_path):
        try:
            with open(stats_path) as f:
                result["pager"] = json.load(f)
        except (OSError, ValueError):
            pass
    # RSS 有沒有落在預算內
    # 「有沒有在預算內」要對**分頁器自己管的東西**判斷,不是對 total RSS。
    #
    # total RSS = 非 expert 權重(常駐,約 267 MiB)
    #          + expert 分頁預算(ST_RAM_BUDGET_MB)
    #          + KV cache(× n_slots)+ compute buffer + 執行檔本身
    # 其中 KV 與 compute 取決於 llama-server 的 slot 數與 ctx,
    # 沒有辦法從這裡準確推估 —— 硬湊一個公式只會得到一個假數字。
    # 所以:total RSS 如實回報(它是量測值),預算判定交給分頁器的帳。
    pg = result.get("pager")
    if pg:
        result["pager_expert_resident_mib"] = round(pg.get("resident_bytes", 0) / MIB, 1)
        result["pager_expert_budget_mib"] = round(pg.get("budget_bytes", 0) / MIB, 1)
        result["pager_expert_within_budget"] = (
            pg.get("resident_bytes", 0) <= pg.get("budget_bytes", 0) * 1.05)
        result["pager_hit_rate"] = pg.get("hit_rate")
    if pg:
        # 分頁器有在管、也能證明有在管的唯一指標
        result["ok"] = result.get("ok", False) and result["pager_expert_within_budget"]
    result["model_size_mib"] = os.path.getsize(a.model) // MIB if os.path.exists(a.model) else None
    ok = "error" not in result and result["peak"]["peak_swap_gb"] == 0
    result["ok"] = ok
    _write(a.out, result)
    print(json.dumps(result, indent=2, ensure_ascii=False))
    return 0 if ok else 1


def _write(path: str, obj: dict) -> None:
    os.makedirs(os.path.dirname(os.path.abspath(path)) or ".", exist_ok=True)
    with open(path, "w") as f:
        json.dump(obj, f, indent=2, ensure_ascii=False)
        f.write("\n")


if __name__ == "__main__":
    sys.exit(main())