mirror of
https://github.com/Helldez/BigMoeOnEdge.git
synced 2026-10-03 03:25:42 +00:00
* feat(moe): --drop-cold-experts, spend quality only where it buys I/O Turbo top-k drops the tail of a routing whether or not those experts were already in RAM. A resident expert costs no flash read, so that trade pays quality for nothing on the ~80% of decode routings that are cache hits. This skips a routed expert only when it is a cache MISS and the router weighted it below frac x (1/top-k). Replayed over the committed route traces at frac 1.0, decode phase, that avoids 66% of flash reads for 9.5% of the router's weight mass, where --n-expert-used 5 avoids 23% for a comparable 10.6% -- about 3x the reads at the same quality cost. Implementation. The decision needs the FINAL router weights, which arrive several nodes after the topk where the streamer normally loads, so load_layer() is deferred to the terminal node of the layer's weight chain. Which node that is depends on the model's gating, so the hook learns it from the graph rather than carrying an architecture table; if it fails to arrive the hook forgets it and re-learns rather than re-betting. A dropped slot has its weight zeroed and its expert id repointed at the routing's top-weighted expert: an unread expert can sit in reserved-but-uncommitted VM and mul_mat_id would touch it anyway, so the kernel is given memory that is certainly resident and multiplies it by exactly zero. Requires the LRU cache -- with --cache-mb 0 residency reads all-miss and the policy would silently degenerate into an unconditional weight cut. Prefill is excluded by default. The top expert is always pinned, so no routing can be emptied at any threshold. Gates: G8a/G8a' prove the deferral and the learned terminal node are transparent (byte-identical output, zero drops, at a threshold below any producible weight); G8b that full strength against a constantly-evicting cache never reaches an unloaded slot; G8c that at top-k 1 dropping is a no-op, pinning both the top-expert guarantee and the threshold tracking the effective top-k. Three existing metrics shift meaning under dropping and the docs now say so: cache_hit_pct rises without the cache serving more (a dropped routing is a miss that is never looked up), and token/layer_demand measure what was staged rather than routed. prefetch.md's "cannot change output" is scoped, limitations.md gains the non-reproducibility entry, and benchmark-method.md warns that reversing the run order cannot distinguish a moved drop rate from a contaminated cell. Off by default in the CLI and in the app. The output is not reproducible -- what gets dropped depends on what the cache held -- so it carries no rows in the README tables, and switching it on by default waits on a published on-device A/B rather than on the replay argument alone. * feat(app): default cache-aware dropping to 75%, measured on device Qwen3.6-35B-A3B (top-k 8 of 256), in-app, cache 3000, one variable changed: 2.549 tok/s off, 3.938 at F=0.75 (+55%), 4.702 at F=1.0 (+84%), with flash reads falling 248 -> 163 -> 48 GiB. Per-token bootstrap intervals separate every pair except off vs 0.50, which overlaps -- at half the uniform share the policy drops 2.7% of routings and buys nothing, which doubles as a negative control that the machinery is free when it does not fire. Run order was 1.0, off, 0.5, 0.75, so the two fastest cells are the first and the LAST; thermal drift would have made the last the worst. The mechanism orders by threshold even though the run order does not. The replay turned out conservative rather than optimistic. It is documented as an upper bound because it cannot model the cache changing in response to dropping: at F=0.75 it was accurate (37% predicted, 34% measured), at F=1.0 it understated (66% predicted, 81% measured). Avoided reads free cache capacity, which raises the hit rate, which leaves fewer misses to drop. 75% rather than 100% is deliberate: it takes the larger part of the win for half the discarded routings (14% against 28%). Quality is still unquantified -- no perplexity number and no side-by-side exists -- so the conservative end of a measured range is the defensible default. The CLI stays off; the byte-identity gates need a deterministic default. Also records cache_hit_pct rising 67.8 -> 90.7% as the documented accounting artefact rather than the cache serving more, and majflt/token as dominated by each run's starting memory state, not by the threshold.
159 lines
6 KiB
Python
159 lines
6 KiB
Python
#!/usr/bin/env python3
|
|
"""Replay a route trace against the cache-aware expert-dropping policy (docs/expert-dropping.md).
|
|
|
|
The policy skips a routed expert when it is a cache MISS and the router weighted it below
|
|
`frac x (1 / n_expert_used)`. A resident expert costs no flash read, so it is never dropped:
|
|
quality is spent only where it buys I/O. This script answers, per threshold, what that trade
|
|
would have been on an already-recorded run:
|
|
|
|
io_saved fraction of MISS BYTES the policy never reads -- the win
|
|
mass_lost fraction of total router weight discarded -- the proxy for the damage
|
|
|
|
The static-k baseline (`--n-expert-used`) is replayed on the same rows, because the only question
|
|
that matters is comparative: at equal io_saved, which policy discards less weight?
|
|
|
|
Two limits, both deliberate:
|
|
* This is a STATIC replay. Skipping a read changes what the cache holds later, so the real hit
|
|
pattern drifts from the recorded one. io_saved is an UPPER BOUND, not a prediction.
|
|
* mass_lost is a proxy. It says how much of the router's mass went away, not what that did to
|
|
the output. Only a quality A/B answers that.
|
|
|
|
A trace recorded with dropping already ON reports its `dropped` column instead of re-deriving it,
|
|
which is how the upper bound above gets checked against a real run.
|
|
|
|
Usage: route-drop-replay.py <route.csv> [<route.csv> ...]
|
|
Stdlib only, like the other analysis scripts here.
|
|
"""
|
|
import os
|
|
import sys
|
|
from collections import defaultdict
|
|
|
|
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
|
from trace_io import read_preamble_csv # noqa: E402
|
|
|
|
MISS = 0
|
|
DECODE, PREFILL = 1, 0
|
|
|
|
|
|
def cells(rows, phase):
|
|
"""Group rows into routing cells: (turn, step, layer) -> [(slot, weight, residency, bytes, dropped)]."""
|
|
out = defaultdict(list)
|
|
for r in rows:
|
|
if int(r["phase"]) != phase:
|
|
continue
|
|
try:
|
|
w = float(r["weight"])
|
|
except ValueError:
|
|
continue # 'nan': the graph exposed no weight node, so no threshold can be applied
|
|
out[(r["turn"], r["step"], r["layer"])].append(
|
|
(int(r["slot"]), w, int(r["residency"]), int(r["expert_bytes"]), int(r.get("dropped", 0) or 0))
|
|
)
|
|
return out
|
|
|
|
|
|
def replay_threshold(cs, thr):
|
|
"""Drop a miss weighted below thr, never the cell's top expert (which the engine also pins)."""
|
|
miss_bytes = dropped_bytes = 0
|
|
total_mass = lost_mass = 0.0
|
|
kept_hist = defaultdict(int)
|
|
for entries in cs.values():
|
|
best = max(range(len(entries)), key=lambda i: entries[i][1])
|
|
kept = 0
|
|
for i, (_slot, w, res, nb, _d) in enumerate(entries):
|
|
total_mass += w
|
|
if res == MISS:
|
|
miss_bytes += nb
|
|
if i != best and w < thr:
|
|
dropped_bytes += nb
|
|
lost_mass += w
|
|
continue
|
|
kept += 1
|
|
kept_hist[kept] += 1
|
|
return miss_bytes, dropped_bytes, total_mass, lost_mass, kept_hist
|
|
|
|
|
|
def replay_static_k(cs, keep_k):
|
|
"""Baseline --n-expert-used: keep the top keep_k slots whatever the cache holds."""
|
|
miss_bytes = dropped_bytes = 0
|
|
total_mass = lost_mass = 0.0
|
|
for entries in cs.values():
|
|
for slot, w, res, nb, _d in entries:
|
|
total_mass += w
|
|
if res == MISS:
|
|
miss_bytes += nb
|
|
if slot >= keep_k:
|
|
lost_mass += w
|
|
if res == MISS:
|
|
dropped_bytes += nb
|
|
return miss_bytes, dropped_bytes, total_mass, lost_mass
|
|
|
|
|
|
def observed(cs):
|
|
"""What a trace recorded with dropping ON actually did. (dropped rows, weight mass, miss bytes)."""
|
|
n_dropped = 0
|
|
lost_mass = total_mass = 0.0
|
|
for entries in cs.values():
|
|
for _slot, w, _res, _nb, d in entries:
|
|
total_mass += w
|
|
if d:
|
|
n_dropped += 1
|
|
lost_mass += w
|
|
return n_dropped, lost_mass, total_mass
|
|
|
|
|
|
def pct(num, den):
|
|
return 100.0 * num / den if den else 0.0
|
|
|
|
|
|
def report(path):
|
|
meta, rows = read_preamble_csv(path)
|
|
k = int(meta.get("n_expert_used", 0) or 0)
|
|
if k <= 0:
|
|
print(f"{path}: no n_expert_used in the preamble; cannot express a threshold")
|
|
return
|
|
print("=" * 78)
|
|
print(f"{os.path.basename(path)} arch={meta.get('arch')} n_expert={meta.get('n_expert')} k={k}")
|
|
print("=" * 78)
|
|
|
|
for phase, label in ((DECODE, "DECODE"), (PREFILL, "PREFILL")):
|
|
cs = cells(rows, phase)
|
|
if not cs:
|
|
print(f"\n[{label}] no rows")
|
|
continue
|
|
n_tot = sum(len(e) for e in cs.values())
|
|
n_miss = sum(1 for e in cs.values() for x in e if x[2] == MISS)
|
|
print(f"\n[{label}] {len(cs)} routing cells, {n_tot} routed experts, {pct(n_miss, n_tot):.1f}% misses")
|
|
|
|
n_drop, lost, total = observed(cs)
|
|
if n_drop:
|
|
print(f" recorded: dropping was ON for this run -- {n_drop} routings dropped "
|
|
f"({pct(n_drop, n_tot):.1f}%), {pct(lost, total):.2f}% of the weight mass")
|
|
|
|
uniform = 1.0 / k
|
|
print(f"\n cache-aware threshold, as a fraction of the uniform share 1/k = {100 * uniform:.2f}%")
|
|
print(f" {'frac':>6} {'thr':>8} {'io_saved':>9} {'mass_lost':>10} surviving experts per cell")
|
|
for frac in (0.25, 0.5, 0.75, 1.0):
|
|
thr = uniform * frac
|
|
mb, db, tm, lm, hist = replay_threshold(cs, thr)
|
|
h = " ".join(f"{kk}:{pct(v, len(cs)):.0f}%" for kk, v in sorted(hist.items()))
|
|
print(f" {frac:6.2f} {100 * thr:7.2f}% {pct(db, mb):8.1f}% {pct(lm, tm):9.2f}% {h}")
|
|
|
|
print(f"\n static-k baseline (--n-expert-used), same rows, for comparison at equal io_saved")
|
|
print(f" {'keep_k':>6} {'io_saved':>9} {'mass_lost':>10}")
|
|
for keep in range(k - 1, 0, -1):
|
|
mb, db, tm, lm = replay_static_k(cs, keep)
|
|
print(f" {keep:6d} {pct(db, mb):8.1f}% {pct(lm, tm):9.2f}%")
|
|
print()
|
|
|
|
|
|
def main(argv):
|
|
if len(argv) < 2:
|
|
print(__doc__)
|
|
return 2
|
|
for p in argv[1:]:
|
|
report(p)
|
|
return 0
|
|
|
|
|
|
if __name__ == "__main__":
|
|
sys.exit(main(sys.argv))
|