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