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
Download tools/ab_test.py from HelloSun/SmallThinker4b: direct link, hf CLI and curl.
- Browser
- Download file 11.4 kB
-
https://huggingface.co/HelloSun/SmallThinker4b/resolve/main/tools/ab_test.py
- Command line
-
hf download hf://HelloSun/SmallThinker4b/tools/ab_test.py
-
curl -L -o ab_test.py https://huggingface.co/HelloSun/SmallThinker4b/resolve/main/tools/ab_test.py
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()) |