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 | """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() | |