#!/usr/bin/env python3 """A/B 驗證:熱參數在 RAM、冷參數到 SSD。 跑三組,每組都在量測前把模型檔的 page cache 丟掉(否則會量到「假 hit」): 1. baseline ST_PAGER=0 上游行為,權重全部常駐 2. budget= ST_PAGER=1 預算 B MiB 熱 expert 在 RAM,冷的丟回 SSD 3. verify 同 2,並額外檢查輸出與 baseline 逐字相同 每組量: - total RSS(**含** file-backed mmap 的權重頁;只看匿名記憶體會嚴重低估) - anon / file 分拆 - VmSwap(必須 0) - /proc//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//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())