File size: 2,988 Bytes
3a7bf50 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 | """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()
|