SmallThinker4b / tools /ab_test.py
Auto Upload Agent
v3: 修 MAP_POPULATE 與 /proc/self/io 解析 bug;A/B 驗證工具;輸出逐字相同已證實
1e54449
Raw History Blame Contribute Delete
11.4 kB
#!/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())