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