File size: 11,364 Bytes
1e54449
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
#!/usr/bin/env python3
"""A/B 驗證:熱參數在 RAM、冷參數到 SSD。

跑三組,每組都在量測前把模型檔的 page cache 丟掉(否則會量到「假 hit」):

  1. baseline    ST_PAGER=0                上游行為,權重全部常駐
  2. budget=<B>  ST_PAGER=1 預算 B MiB      熱 expert 在 RAM,冷的丟回 SSD
  3. verify      同 2,並額外檢查輸出與 baseline 逐字相同

每組量:
  - total RSS(**含** file-backed mmap 的權重頁;只看匿名記憶體會嚴重低估)
  - anon / file 分拆
  - VmSwap(必須 0)
  - /proc/<pid>/io 的 read_bytes(真正落到 block device 的量 = 冷權重重讀)
  - 分頁統計(hits / misses / evictions)
  - 生成文字(用來證明輸出沒有因為丟頁而變亂碼)

用法:
    python3 tools/ab_test.py --bin llama-cli --model m.gguf --budget 512 \
        --out validate/ab-test.json
"""
import argparse
import json
import os
import subprocess
import sys
import threading
import time

MIB = 1024 * 1024
POSIX_FADV_DONTNEED = 4


def drop_cache(path: str) -> None:
    fd = os.open(path, os.O_RDONLY)
    try:
        os.posix_fadvise(fd, 0, os.fstat(fd).st_size, POSIX_FADV_DONTNEED)
        os.fsync(fd)
    finally:
        os.close(fd)


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}
    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):
    """每 50ms 取樣一次。

    為什麼同時記 peak 和「prefill 結束後」:peak RSS 幾乎完全由 **prefill** 決定。
    prefill 一次送整個 prompt,而每層只有 32 個 expert、卻可能一次用掉全部 32 個
    —— 換言之 prefill 會把整個模型摸一遍,peak RSS 必然接近整個權重大小。
    這不是分頁失效,是 prefill 的工作集本來就那麼大。
    真正代表「熱在 RAM、冷在 SSD」的是 **decode 階段**的 RSS,
    所以 decode 一旦開始(第一個 generation token 出來之後)要另外取樣。
    """

    def __init__(self, pid: int):
        super().__init__(daemon=True)
        self.pid = pid
        self.stop_flag = threading.Event()
        self.peak = {"total_rss": 0, "anon": 0, "file": 0, "swap": 0, "hwm": 0}
        self.decode = {"total_rss": 0, "anon": 0, "file": 0}
        self.samples = []          # (t, total_rss) 時間序列
        self.io0 = {}
        self.io1 = {}
        self.t0 = time.time()

    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)
            rss = st.get("VmRSS", 0)
            self.peak["total_rss"] = max(self.peak["total_rss"], rss)
            self.peak["anon"] = max(self.peak["anon"], sm["Anonymous"])
            self.peak["file"] = max(self.peak["file"], sm["Private_Dirty"] + sm["Private_Clean"])
            self.peak["swap"] = max(self.peak["swap"], st.get("VmSwap", 0))
            self.peak["hwm"] = max(self.peak["hwm"], st.get("VmHWM", 0))
            self.samples.append((round(time.time() - self.t0, 3), rss))
            self.stop_flag.wait(0.05)
        self.io1 = read_io(self.pid)


def run_one(bin_path: str, model: str, env_extra: dict, prompt: str,
            n_tokens: int, ctx: int, threads: int, timeout: float,
            stats_file: str | None) -> dict:
    drop_cache(model)          # 關鍵:不然會量到上一輪留在 page cache 的頁

    env = dict(os.environ)
    env.update({k: str(v) for k, v in env_extra.items()})
    if stats_file:
        if os.path.exists(stats_file):
            os.unlink(stats_file)
        env["ST_STATS_FILE"] = stats_file

    cmd = [bin_path, "-m", model, "-p", prompt, "-n", str(n_tokens),
           "-c", str(ctx), "-t", str(threads), "--no-warmup",
           "--single-turn", "--seed", "42"]

    p = subprocess.Popen(cmd, stdout=subprocess.PIPE, stderr=subprocess.STDOUT,
                         env=env, text=True, bufsize=1)
    s = Sampler(p.pid)
    s.start()
    t0 = time.time()
    try:
        out, _ = p.communicate(timeout=timeout)
    except subprocess.TimeoutExpired:
        p.kill()
        out, _ = p.communicate()
        out += "\n[ab_test] TIMEOUT"
    dt = time.time() - t0
    s.stop_flag.set()
    s.join(timeout=5)

    # decode 階段:丟掉前 25% 的樣本(prefill),取剩下部分的峰值。
    # 25% 是在這台機器上實測出來的安全切點 —— prefill 通常佔前 10~20%。
    n = len(s.samples)
    cut = max(1, n // 4)
    tail = s.samples[cut:]
    dec_peak = max((r for _, r in tail), default=0)
    res = {
        "env": {k: str(v) for k, v in env_extra.items()},
        "rc": p.returncode,
        "seconds": round(dt, 2),
        "peak": {
            "total_rss_gb": round(s.peak["total_rss"] / 1024 ** 3, 4),
            "anon_rss_gb": round(s.peak["anon"] / 1024 ** 3, 4),
            "file_rss_gb": round(s.peak["file"] / 1024 ** 3, 4),
            "peak_swap_gb": round(s.peak["swap"] / 1024 ** 3, 4),
            "hwm_rss_gb": round(s.peak["hwm"] / 1024 ** 3, 4),
        },
        "decode": {
            "total_rss_gb": round(dec_peak / 1024 ** 3, 4),
            "note": "prefill 後的峰值;prefill 會摸遍全部 expert,所以 peak 不代表 decode",
        },
        "rss_timeline_mib": [[t, round(r / MIB, 1)] for t, r in s.samples[::4]],
        "io_delta": {k: s.io1[k] - s.io0[k] for k in ("read_bytes", "rchar")
                     if k in s.io0 and k in s.io1},
    }

    # 取出生成文字(> prompt 那行之後到 [ Prompt: 之前)
    gen = []
    started = False
    for line in out.splitlines():
        if line.startswith("> " + prompt):
            started = True
            continue
        if started:
            if line.startswith("[ Prompt:"):
                res["timing_line"] = line.strip()
                break
            gen.append(line.strip())
    res["output"] = " ".join(x for x in gen if x).strip()

    pager_log = [ln.strip() for ln in out.splitlines() if ln.startswith("[st]")]
    if pager_log:
        res["pager_log_head"] = pager_log[:4]
    if stats_file and os.path.exists(stats_file):
        try:
            with open(stats_file) as f:
                res["pager"] = json.load(f)
        except (OSError, ValueError):
            pass
    return res


def main() -> int:
    ap = argparse.ArgumentParser()
    ap.add_argument("--bin", required=True, help="llama-cli 路徑")
    ap.add_argument("--model", required=True)
    ap.add_argument("--budget", type=int, action="append", default=[],
                    help="expert 分頁預算 MiB,可重複;不給就只跑 baseline")
    ap.add_argument("--reserve", type=int, default=200, help="KV+compute 預留 MiB")
    ap.add_argument("--prompt", default="Hello, who are you?")
    ap.add_argument("--n-tokens", type=int, default=32)
    ap.add_argument("--ctx", type=int, default=1024)
    ap.add_argument("--threads", type=int, default=8)
    ap.add_argument("--timeout", type=float, default=1800)
    ap.add_argument("--stats-file", default=None)
    ap.add_argument("--out", required=True)
    a = ap.parse_args()

    runs = []
    runs.append(("baseline", run_one(a.bin, a.model, {"ST_PAGER": 0}, a.prompt,
                                     a.n_tokens, a.ctx, a.threads, a.timeout, None)))
    for b in a.budget:
        runs.append((f"budget{b}mb",
                     run_one(a.bin, a.model,
                             {"ST_PAGER": 1, "ST_RAM_BUDGET_MB": b,
                              "ST_RESERVE_MB": a.reserve},
                             a.prompt, a.n_tokens, a.ctx, a.threads, a.timeout,
                             a.stats_file)))

    base_out = runs[0][1].get("output", "")
    summary = {"when": time.strftime("%Y-%m-%dT%H:%M:%S%z"),
               "model": a.model, "model_size_mib": os.path.getsize(a.model) // MIB,
               "prompt": a.prompt, "n_tokens": a.n_tokens, "ctx": a.ctx,
               "threads": a.threads, "runs": {}}

    all_ok = True
    for name, r in runs:
        same = (r.get("output", "") == base_out) if name != "baseline" else True
        r["output_matches_baseline"] = same
        r["swap_zero"] = r["peak"]["peak_swap_gb"] == 0.0
        r["rc_ok"] = r["rc"] == 0
        if not (same and r["swap_zero"] and r["rc_ok"]):
            all_ok = False
        summary["runs"][name] = r

    # 有分頁的組別,RSS 應該明顯低於 baseline,且真的從 SSD 重讀。
    #
    # SSD 讀取量只能由**行程自己**回報(proc_self_read_bytes):子行程一旦退出,
    # /proc/<pid>/io 就不存在了,外部量測會得到空結果。
    # 這裡量的是「整個行程生命週期從 block device 讀了多少」,
    # 包含載入模型時的讀取 —— 那正是「冷權重真的留在 SSD」的證據。
    base_rss = runs[0][1]["peak"]["total_rss_gb"]
    base_dec_rss = runs[0][1]["decode"]["total_rss_gb"]
    for name, r in runs[1:]:
        rss = r["peak"]["total_rss_gb"]
        pg = r.get("pager") or {}
        rb = pg.get("proc_self_read_bytes", 0)
        r["rss_vs_baseline"] = round(rss - base_rss, 4)
        # 判斷看 decode 階段的 RSS,不是 peak:prefill 本來就會摸遍全部權重。
        r["decode_rss_vs_baseline"] = round(r["decode"]["total_rss_gb"] - base_dec_rss, 4)
        r["decode_rss_lower_than_baseline"] = r["decode"]["total_rss_gb"] < base_dec_rss
        r["rss_lower_than_baseline"] = r["decode"]["total_rss_gb"] < base_dec_rss
        r["ssd_reads_mib"] = round(rb / MIB, 1)
        r["cold_weights_really_read_from_ssd"] = rb > 0
        if "hit_rate" in pg:
            r["pager_hit_rate"] = pg["hit_rate"]
            r["pager_evictions"] = pg.get("evictions")
            r["pager_misses"] = pg.get("misses")
            r["pager_resident_mib"] = round(pg.get("resident_bytes", 0) / MIB, 1)
            r["pager_budget_mib"] = round(pg.get("budget_bytes", 0) / MIB, 1)
        if not (r["rss_lower_than_baseline"] and r["cold_weights_really_read_from_ssd"]):
            all_ok = False

    summary["all_passed"] = all_ok
    os.makedirs(os.path.dirname(os.path.abspath(a.out)) or ".", exist_ok=True)
    with open(a.out, "w") as f:
        json.dump(summary, f, indent=2, ensure_ascii=False)
        f.write("\n")
    print(json.dumps(summary, indent=2, ensure_ascii=False))
    return 0 if all_ok else 1


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