onw / sim_expert_cache.py
ryugyosoft's picture
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
Raw
History Blame Contribute Delete
2.99 kB
"""Replay an expert trace (trace_experts.py) against expert caches of several sizes: hit rate, experts read per
decoded token / per prompt block, and the decode / prefill time they add at a given SSD speed.
usage: python sim_expert_cache.py TRACE.json [GB/s=3.3] [MB per expert=1.58]"""
import json, sys
from collections import OrderedDict
def simulate(trace, cap, policy="lru"):
cache, freq = OrderedDict(), {}
dec_reads = dec_steps = pre_reads = pre_blocks = hits = total = 0
for layer, S, n, experts in trace:
need = [(layer, e) for e in experts]
miss = 0
for k in need:
total += 1
freq[k] = freq.get(k, 0) + 1
if k in cache:
hits += 1
cache.move_to_end(k)
else:
miss += 1
pinned = set(need)
while len(cache) >= cap:
if policy == "lfu": # evict the least used (ties: least recent) outside this layer's set
victim = min((x for x in cache if x not in pinned), key=lambda x: (freq[x], 0))
else:
victim = next(x for x in cache if x not in pinned)
del cache[victim]
cache[k] = True
if S == 1:
dec_reads += miss
if layer == 0:
dec_steps += 1
else:
pre_reads += miss
if layer == 0:
pre_blocks += 1
return {"hit": hits / max(total, 1), "per_token": dec_reads / max(dec_steps, 1),
"per_block": pre_reads / max(pre_blocks, 1), "tokens": dec_steps, "blocks": pre_blocks}
def main():
t = json.load(open(sys.argv[1]))
gbps = float(sys.argv[2]) if len(sys.argv) > 2 else 3.3
mb = float(sys.argv[3]) if len(sys.argv) > 3 else 1.58
trace = t["trace"]
layers = len({x[0] for x in trace})
first = min(x[0] for x in trace)
trace = [[x[0] - first, *x[1:]] for x in trace]
E = t["E"]
print(f"{len(trace)} binds, {layers} MoE layers x {E} experts = {layers * E} experts ({layers * E * mb / 1024:.1f} GB)")
for run in t["runs"]:
print(f" {run['name']:10s} prompt {run['prompt_tokens']:4d} (reused {run['cached']}), {run['completion']} tokens")
print(f"\nSSD {gbps} GB/s, {mb} MB per expert")
print("cache GB share policy hit% reads/token +ms/token reads/16-token block +s/block")
for gb in (2, 3, 4, 5, 6, 8, 10):
cap = int(gb * 1024 / mb)
if cap >= layers * E:
continue
for pol in ("lru", "lfu"):
r = simulate(trace, cap, pol)
ms = r["per_token"] * mb / 1024 / gbps * 1000
sb = r["per_block"] * mb / 1024 / gbps
print(f"{gb:6d} {cap / (layers * E) * 100:4.0f}% {pol} {r['hit'] * 100:5.1f} {r['per_token']:9.1f} "
f"{ms:7.0f} {r['per_block']:9.1f} {sb:5.2f}")
if __name__ == "__main__":
main()