sparkle/mlx_streaming/runtime/run_mmap.py
fiser_jun 4745f264b2
2026-08-04 14:34:00 +08:00

70 lines
2.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""路线 Alazy+ 若支持则 mmap加载可选设常驻/缓存上限,量内存与速度。
环境探针结论mlx-lm 0.31.3load 支持 lazy但不支持 use_mmap。
本脚本用 _filtered_load_kwargs 只传该版本真正支持的参数,缺失的自动剔除。
"""
import os
import time
import json
import inspect
import mlx.core as mx
from mlx_lm import load, generate
from mlx_streaming.core.mem import snapshot, reset_peak, clear_cache
MODEL = os.environ.get("MODEL", "models/Qwen3-Next-80B-A3B-Instruct-MLX-8bit")
PROMPT = os.environ.get("PROMPT", "用三句话解释什么是混合专家模型。")
MAXTOK = int(os.environ.get("MAXTOK", "128"))
WIRED_GB = os.environ.get("WIRED_GB") # 例如 "6";空=不设
CACHE_GB = os.environ.get("CACHE_GB") # 例如 "1";空=不设
def _maybe(fn_name, val_bytes):
fn = getattr(mx, fn_name, None) or getattr(getattr(mx, "metal", object()), fn_name, None)
if fn and val_bytes is not None:
prev = fn(int(val_bytes))
print(f"{fn_name}({int(val_bytes)}) 旧值={prev}")
def _filtered_load_kwargs():
sig = inspect.signature(load)
want = {"lazy": True, "use_mmap": True}
return {k: v for k, v in want.items() if k in sig.parameters}
def main():
if WIRED_GB:
_maybe("set_wired_limit", float(WIRED_GB) * 1e9)
if CACHE_GB:
_maybe("set_cache_limit", float(CACHE_GB) * 1e9)
reset_peak()
kw = _filtered_load_kwargs()
print("load kwargs:", kw)
t0 = time.perf_counter()
model, tok = load(MODEL, **kw) # 不强制 eval 全部参数!
t1 = time.perf_counter()
after_load = snapshot() # 关键:此时 RSS 应远低于基线(专家还没换入)
text = generate(model, tok, prompt=PROMPT, max_tokens=MAXTOK, verbose=False)
t2 = time.perf_counter()
clear_cache()
after_gen = snapshot()
out = {
"mode": "mmap_lazy", "model": MODEL,
"wired_gb": WIRED_GB, "cache_gb": CACHE_GB,
"load_s": round(t1 - t0, 2), "gen_s": round(t2 - t1, 2),
"tok_per_s": round(MAXTOK / (t2 - t1), 2),
"rss_gb_after_load": round(after_load.rss_bytes / 1e9, 2),
"rss_gb_after_gen": round(after_gen.rss_bytes / 1e9, 2),
"mlx_active_gb_after_gen": round(after_gen.mlx_active_bytes / 1e9, 2),
"mlx_peak_gb": round(after_gen.mlx_peak_bytes / 1e9, 2),
}
print(json.dumps(out, ensure_ascii=False, indent=2))
if __name__ == "__main__":
main()