BigMoeOnEdge/scripts/route-drop-replay.py
Helldez bea5a0b99e
feat(moe): --drop-cold-experts — spend quality only where it buys I/O (#95)
* 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.
2026-07-22 17:21:55 +02:00

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