v0.6.0: stream MoE experts from the SSD (LRU cache, direct I/O) when the model does not fit; expert memory setting
3a7bf50 verified Download test_stream.py from ryugyosoft/onw: direct link, hf CLI and curl.
- Browser
- Download file 2.16 kB
-
https://huggingface.co/ryugyosoft/onw/resolve/main/test_stream.py
- Command line
-
hf download hf://ryugyosoft/onw/test_stream.py
-
curl -L -o test_stream.py https://huggingface.co/ryugyosoft/onw/resolve/main/test_stream.py
2.16 kB
| """Expert streaming check: answers, decode speed, expert-cache hit rate and process memory, with ONW_EXPERT_GB set | |
| (experts beyond that stay on the SSD) - compare with a run without it. | |
| usage: ONW_EXPERT_GB=5 python test_stream.py MODEL_DIR [IMAGE]""" | |
| import os, sys, time | |
| import psutil | |
| from onw.chat import ChatEngine | |
| Q = ["NPUとGPUの違いを、身近なたとえを使って説明してください。", | |
| "ある商品を定価の2割引きで買うと960円でした。定価はいくらですか?途中の式も書いてください。"] | |
| def run(e, msgs, n=160): | |
| parts, st = [], None | |
| for x in e.stream_chat([dict(m) for m in msgs], n): | |
| if isinstance(x, dict): | |
| st = x | |
| else: | |
| parts.append(x) | |
| return "".join(parts), st | |
| def main(): | |
| t0 = time.time() | |
| e = ChatEngine(sys.argv[1], "NPU", pld=False) | |
| print(f"load {time.time() - t0:.0f}s, ONW_EXPERT_GB={os.environ.get('ONW_EXPERT_GB')}", flush=True) | |
| bank = e.model.bank | |
| qs = list(Q) | |
| if len(sys.argv) > 2: | |
| from PIL import Image | |
| qs.insert(1, [{"type": "image", "image": Image.open(sys.argv[2])}, {"type": "text", "text": "この画像について説明してください。"}]) | |
| for q in qs: | |
| e.checkpoint = None | |
| before = bank.stats() if hasattr(bank, "stats") else None | |
| text, st = run(e, [{"role": "user", "content": q}]) | |
| s = bank.stats() if hasattr(bank, "stats") else None | |
| extra = "" | |
| if s: | |
| h = s["hit_rate"] | |
| extra = (f" | cache hit {h * 100:.1f}% (cumulative), read {s['read_gb'] - before['read_gb']:.2f} GB in " | |
| f"{s['read_s'] - before['read_s']:.1f}s") | |
| print(f"prompt {st['prompt_tokens']} tok in {st['prefill_ms'] / 1000:.1f}s, decode {st['decode_tok_s']:.2f} tok/s{extra}\n" | |
| f" {text[:80]!r}", flush=True) | |
| m = psutil.Process().memory_info() | |
| print(f"MEM private {getattr(m, 'private', m.rss) / 2**30:.1f} GB, working set {m.rss / 2**30:.1f} GB, " | |
| f"peak working set {getattr(m, 'peak_wset', 0) / 2**30:.1f} GB", flush=True) | |
| if __name__ == "__main__": | |
| main() | |