Text Generation
llama-cpp-python
GGUF
llama.cpp
Mixture of Experts
ssd-offload
smallthinker
expert-paging
low-ram
Instructions to use HelloSun/SmallThinker4b with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- llama-cpp-python
How to use HelloSun/SmallThinker4b with llama-cpp-python:
# !pip install llama-cpp-python from llama_cpp import Llama llm = Llama.from_pretrained( repo_id="HelloSun/SmallThinker4b", filename="{{GGUF_FILE}}", )output = llm( "Once upon a time,", max_tokens=512, echo=True ) print(output)
- Notebooks
- Google Colab
- Kaggle
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()) |