sparkle/mlx_streaming/runtime/run_mtp_spec.py
fiser_jun 4745f264b2
2026-08-04 14:34:00 +08:00

410 lines
18 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""真实 80B MTP 自投机基准 + 与非投机贪婪逐 token 一致性校验。
环境变量:K / MAXTOK / PROMPT / QN_CONFIG / MTP_OUT(其余主模型路径见 validate_mtp)。
"""
import json
import os
import statistics
import time
import mlx.core as mx
from mlx_lm.models.qwen3_next import ModelArgs
from mlx_streaming.core.moe import native_moe
from mlx_streaming.core.mem import snapshot, reset_peak
from mlx_streaming.mtp.drafter import MTPDrafter
from mlx_streaming.mtp.generate import forward_with_hidden, mtp_generate, prefill_chunked
from mlx_streaming.mtp.qwen3_next_mtp import load_mtp
from mlx_streaming.model_builder import build_streaming_model
from mlx_streaming import config as _cfg
QN_CONFIG = _cfg.qn_config()
MTP_OUT = _cfg.mtp_out()
PROMPT = os.environ.get("PROMPT", "用三句话解释什么是混合专家模型。")
MAXTOK = int(os.environ.get("MAXTOK", "96"))
K = int(os.environ.get("K", "3"))
PIN_HOT = int(os.environ.get("PIN_HOT", "0"))
PIN_CAL_TOK = int(os.environ.get("PIN_CAL_TOK", "32"))
# 稳态测速:warmup 跑满 MAXTOK(热 baseline+spec 两条路径的 Metal kernel 与常驻专家池/预取),
# 再各重复 REPEAT 次取中位数,避免冷启动编译/补池污染绝对 tok/s。WARMUP_TOK=0 关 warmup。
WARMUP_TOK = int(os.environ.get("WARMUP_TOK", str(MAXTOK)))
REPEAT = int(os.environ.get("REPEAT", "3"))
def _spec_once(model, drafter, tok, n, k):
ids, stats = mtp_generate(model, drafter, tok,
mx.array([tok.encode(PROMPT)]),
n, K=k, ids_mode=True, profile=True)
tps = round(stats["tokens"] / stats["wall_s"], 2)
return ids, stats, tps
def _baseline_greedy(model, tok, prompt, n):
cache = model.make_cache()
ids = mx.array([tok.encode(prompt)])
t0 = time.perf_counter()
# prefill 分块:把整段 prompt 的激活峰值压到与 decode 同稳态(见 config.prefill_chunk)。
logits, _ = prefill_chunked(model, ids, cache)
_dump_margin = bool(os.environ.get("DUMP_MARGIN"))
out = []
for _ in range(n):
lg = logits[:, -1, :]
nxt = int(mx.argmax(lg))
out.append(nxt)
if _dump_margin:
# 诊断:每步 top-2 logit 与差值,用于判定路径间发散是「FP 近平局」还是「真错」。
v = lg.reshape(-1)
top2 = mx.argpartition(-v, 2)[:2]
t2 = [int(x) for x in top2.tolist()]
vv = {t: float(v[t]) for t in t2}
st = sorted(t2, key=lambda t: -vv[t])
print(f"MARGIN step={len(out)-1} top1={st[0]}({vv[st[0]]:.5f}) "
f"top2={st[1]}({vv[st[1]]:.5f}) gap={vv[st[0]]-vv[st[1]]:.6f}", flush=True)
cur = mx.array([[nxt]])
mx.eval(cur)
if len(out) >= n:
break
logits, _ = forward_with_hidden(model, cur, cache)
return out, round(n / (time.perf_counter() - t0), 2)
def main():
reset_peak()
model, tok, store = build_streaming_model()
with open(QN_CONFIG) as f:
args = ModelArgs.from_dict(json.load(f))
mtp = load_mtp(args, MTP_OUT, quantize=True)
mtp.embed_tokens = model.model.embed_tokens # 共享主模型 embedding
drafter = MTPDrafter(mtp, model.lm_head)
# warmup:同时跑 baseline + spec 两条路径,编译 Metal kernel(含 multistate/batch verify)
# 并把专家常驻池/预取热起来,确保后续测的是稳态而非冷启动。默认 warmup=MAXTOK 跑满全长。
if WARMUP_TOK > 0:
_baseline_greedy(model, tok, PROMPT, WARMUP_TOK)
_spec_once(model, drafter, tok, WARMUP_TOK, K)
# ---- baseline 稳态:重复 REPEAT 次取中位数;最后一次清零统计供命中率口径 ----
base = None
base_tps_runs = []
for r in range(REPEAT):
if r == REPEAT - 1:
store.reset_stats()
base, bt = _baseline_greedy(model, tok, PROMPT, MAXTOK)
base_tps_runs.append(bt)
base_tps = statistics.median(base_tps_runs)
base_miss, base_hit = store.misses, store.hits
base_prefetch_loads = store._resident.prefetch_loads
base_prefetch_hits = store._resident.prefetch_hits
# 可选:baseline 之后、spec 之前校准每层热专家,并预取/钉入 resident pool。
# 这样 disk_load_ratio 的分母仍是未 pin baseline能直接衡量 pin 是否压低 spec miss。
if PIN_HOT > 0:
store.record = True
_baseline_greedy(model, tok, PROMPT, PIN_CAL_TOK)
for li in store.recorded_layers():
store.pin(li, store.hot(li, PIN_HOT))
store.record = False
store.reset_stats()
# 诊断:开 handler 触发时刻探针(仅覆盖 spec 阶段,enable 会清零旧日志)。
_hprof = bool(os.environ.get("STAGING_HPROF"))
if _hprof:
from mlx_streaming import native_moe_ext as _Nhp
_Nhp.staging_hprof_enable(True)
# ---- spec 稳态:重复 REPEAT 次取中位数;最后一次清零统计 + reset_peak + 保留其 stats ----
ids, stats, spec_tps_runs = None, None, []
for r in range(REPEAT):
if r == REPEAT - 1:
store.reset_stats()
reset_peak()
from mlx_streaming.core.profiling import tprof_reset, union_reset
tprof_reset() # 探针只统计最终测量轮(与命中率/内存口径一致)
union_reset() # 并集专家数也只统计最终轮
ids, stats, st = _spec_once(model, drafter, tok, MAXTOK, K)
spec_tps_runs.append(st)
spec_tps = statistics.median(spec_tps_runs)
spec_miss, spec_hit = store.misses, store.hits
spec_prefetch_loads = store._resident.prefetch_loads
spec_prefetch_hits = store._resident.prefetch_hits
after = snapshot()
proj = stats.get("proj_no_replay_tps", 0.0)
result = {
"K": K,
"max_tokens": MAXTOK,
"warmup_tok": WARMUP_TOK,
"repeat": REPEAT,
"exact_match": ids == base,
"n_mismatch": sum(1 for a, b in zip(ids, base) if a != b),
"avg_accept_len": stats["avg_accept_len"],
"steps": stats["steps"],
"verify_mode": stats.get("verify_mode"),
"direct_commits": stats.get("direct_commits"),
"fallback_replays": stats.get("fallback_replays"),
"replayed_tokens": stats.get("replayed_tokens"),
"spec_tok_per_s": spec_tps,
"baseline_tok_per_s": base_tps,
"speedup": round(spec_tps / max(base_tps, 1e-6), 2),
"spec_tps_runs": spec_tps_runs,
"baseline_tps_runs": base_tps_runs,
"spec_tps_minmax": [min(spec_tps_runs), max(spec_tps_runs)],
# 分段计时与「重放免费」投机上限
"t_draft_s": stats.get("t_draft_s"),
"t_snap_s": stats.get("t_snap_s"),
"t_verify_s": stats.get("t_verify_s"),
"t_commit_s": stats.get("t_commit_s"),
"t_replay_s": stats.get("t_replay_s"),
"t_sync_s": stats.get("t_sync_s"),
"t_finalize_s": stats.get("t_finalize_s"),
"proj_no_replay_tps": proj,
"proj_no_replay_speedup": round(proj / max(base_tps, 1e-6), 2),
"baseline_disk_loads": base_miss,
"baseline_prefetch_loads": base_prefetch_loads,
"baseline_prefetch_hits": base_prefetch_hits,
"baseline_hit_rate": round(base_hit / max(base_hit + base_miss, 1), 3),
"spec_disk_loads": spec_miss,
"spec_prefetch_loads": spec_prefetch_loads,
"spec_prefetch_hits": spec_prefetch_hits,
"spec_hit_rate": round(spec_hit / max(spec_hit + spec_miss, 1), 3),
"disk_load_ratio": round(spec_miss / max(base_miss, 1), 2),
# 双源 acquire 分路计数:n_miss==0 走全 GPU 快路径;任一路由 miss 则整层落 host 慢路径
# (.tolist 全批同步 + demand 读盘)。fallback 占比高 → 即使 hit 高,慢路径仍按"层"频繁触发。
"gpu_fastpath": getattr(store._resident, "gpu_fastpath", None),
"gpu_fallback": getattr(store._resident, "gpu_fallback", None),
"pin_hot": PIN_HOT,
"pin_cal_tok": PIN_CAL_TOK,
"pinned_experts": store.pinned_count(),
"expert_slots": store.capacity,
"mlx_active_gb": round(after.mlx_active_bytes / 1e9, 2),
"mlx_peak_gb": round(after.mlx_peak_bytes / 1e9, 2),
"rss_gb": round(after.rss_bytes / 1e9, 2),
# 内存分块(清缓冲后的真实常驻):A 权重 / B 专家池 / C staging / D MTP / F 激活
"mem_breakdown": _mem_breakdown(model, store, mtp),
"prefill_chunk": _cfg.prefill_chunk(),
"native_stage_cache": native_moe.stage_cache_stats(),
"bg_stats": (store._bg.stats() if getattr(store, "_bg", None) is not None else None),
"window_prof": _window_prof(),
"predict_recall": _predict_recall(),
"miss_attrib": _miss_attrib(),
"prefetch_tprof": _prefetch_tprof(stats.get("wall_s")),
"union_experts": _union_prof(),
}
# 噪声地板测量口径:DUMP_IDS=1 时把 baseline greedy 与 spec 的完整 token 序列打进日志,
# 供跨进程 run-to-run 逐位对比(默认关闭,不污染常规输出)。
if os.environ.get("DUMP_IDS"):
print("DUMP_BASE_IDS " + json.dumps(list(base)))
print("DUMP_SPEC_IDS " + json.dumps(list(ids)))
# 字节真值校验自证口径:开 STG_VERIFY 时,把两处校验器的累计计数
# (ok/bad/calls)打进日志。关键:让「0 BAD」可判真伪——若 calls==0 说明本配置根本
# 没触发该校验器(如 STG_VERIFY 在 zerocopy_dual 路径不接线),此时 0 BAD 是空结论。
if os.environ.get("STG_VERIFY"):
from mlx_streaming.core.cache import resident_pool as _rp_mod
from mlx_streaming.core.cache import virtual_pool as _vp_mod
_vsum = {
"STG_VERIFY.resident(verify_acquire_bytes)": dict(_rp_mod._stg_verify_state),
"STG_VERIFY.virtual(_verify_native_bytes)": dict(_vp_mod._stg_verify_state),
}
print("VERIFY_SUMMARY " + json.dumps(_vsum, ensure_ascii=False))
print(json.dumps(result, ensure_ascii=False, indent=2))
if _hprof:
# dump (gen, layer, t_fire) 原始日志,供离线分析回调触发时刻分布。
_flat = _Nhp.staging_hprof_get() # 扁平 [gen,layer,t, ...]
_path = os.environ.get("STAGING_HPROF_OUT", "/tmp/ab/hprof.jsonl")
os.makedirs(os.path.dirname(_path), exist_ok=True)
_n = 0
with open(_path, "w") as _f:
for i in range(0, len(_flat), 3):
_f.write(json.dumps([int(_flat[i]), int(_flat[i + 1]), float(_flat[i + 2])]) + "\n")
_n += 1
print(f"[hprof] wrote {_n} handler records to {_path}")
def _tree_nbytes(obj):
from mlx.utils import tree_flatten
return sum(v.nbytes for _, v in tree_flatten(obj) if isinstance(v, mx.array))
def _mem_breakdown(model, store, mtp):
"""把 decode 稳态内存拆成各块(清掉 MLX 可回收缓冲后量真实常驻)。
A 常驻非专家权重 / B 专家常驻池 / C staging 预取 / D MTP drafter(扣共享 embedding)/
F 激活+临时(= active底 ABCD)。返回 GiB 字典。
"""
GIB = 1024 ** 3
A = _tree_nbytes(model.parameters())
rp = store._resident
B = sum(v.nbytes for pool in rp._pools.values() for v in pool.values())
C = 0
stg = getattr(store, "_staging", None)
if stg is not None:
from mlx.utils import tree_flatten
C = sum(v.nbytes for _, v in tree_flatten(getattr(stg, "__dict__", {}))
if isinstance(v, mx.array))
# D:MTP drafter 全部权重减去与主模型共享的 embedding(run 里 mtp.embed_tokens = 主模型的)
D = _tree_nbytes(mtp.parameters())
emb = getattr(getattr(mtp, "embed_tokens", None), "weight", None)
shared = 0
if isinstance(emb, mx.array):
# 共享 embedding 在 A 已计入;若 mtp 的 embed 与主模型同一对象则从 D 扣除避免重复
shared = _tree_nbytes(mtp.embed_tokens) if hasattr(mtp, "embed_tokens") else 0
D = max(0, D - shared)
mx.clear_cache()
active = mx.get_active_memory()
peak = mx.get_peak_memory()
F = max(0, active - (A + B + C + D))
pool_rows = sum(rp.allocated_slots(l) for l in rp._pools)
return {
"A_weights_gib": round(A / GIB, 3),
"B_expert_pool_gib": round(B / GIB, 3),
"B_pool_rows": pool_rows,
"C_staging_gib": round(C / GIB, 3),
"D_mtp_drafter_gib": round(D / GIB, 3),
"F_activation_temp_gib": round(F / GIB, 3),
"active_live_gib": round(active / GIB, 3),
"peak_gib": round(peak / GIB, 3),
}
def _union_prof():
"""按 seq 分桶汇报每层路由专家并集大小(avg/max/p99/分层);seq=K 桶即 MTP verify 的专家并集。
需 UNION_PROF=1 才有数据。返回 {seq: {avg_union, max_union, p99_union, per_layer...}};
verify 桶 = seq==K(verify_in=[x, d_1..d_{K-1}] 恰 K 个 token,见 mtp/generate.py)。
U_max 即该桶各层并集的全局最大值,是池 cap 下限的正确性依据(cap 必须 ≥ U_max 才保证不溢出)。
"""
from mlx_streaming.core.profiling import UNION_PROF as U, UNION_SAMPLES as S
if not U:
return None
def _p(vals, q):
# 最近秩法分位数(vals 已升序):q∈[0,1]。样本少时直接给保守上界。
if not vals:
return None
idx = max(0, min(len(vals) - 1, int(round(q * (len(vals) - 1)))))
return vals[idx]
out = {}
for seq in sorted(U):
s, n = U[seq]
rec = {"avg_union": round(s / max(1, n), 2), "n_layer_calls": n}
# 由分层原始样本汇总出该 seq 桶的 U_max / p99cap 下限的正确性依据)。
layer_map = S.get(seq)
if layer_map:
flat = sorted(v for lst in layer_map.values() for v in lst)
rec["max_union"] = flat[-1]
rec["p99_union"] = _p(flat, 0.99)
rec["p50_union"] = _p(flat, 0.50)
rec["n_samples"] = len(flat)
# 分层分布:每层的 max / avg / 样本数(按层号排序)。
per_layer = {}
for li in sorted(layer_map):
lst = layer_map[li]
per_layer[li] = {
"max": max(lst),
"avg": round(sum(lst) / len(lst), 2),
"n": len(lst),
}
rec["per_layer"] = per_layer
rec["max_layer_idx"] = max(per_layer, key=lambda li: per_layer[li]["max"])
out[f"seq{seq}"] = rec
# 便捷:verify 桶 = seq==K(verify_in=[x, d_1..d_{K-1}] 恰 K 个 token)
verify_seq = K
verify = out.get(f"seq{verify_seq}")
return {"by_seq": out, "verify_seq": verify_seq,
"verify_avg_union": (verify or {}).get("avg_union"),
"verify_max_union": (verify or {}).get("max_union"),
"verify_p99_union": (verify or {}).get("p99_union")}
def _miss_attrib():
from mlx_streaming.core.profiling import MISS_ATTRIB as M
if not M["n"]:
return None
routed = max(1, M["routed"])
miss = M["miss_A_predicted"] + M["miss_B_unpredicted"]
out = {
"routed": M["routed"],
"hit_rate": round(M["resident_hit"] / routed, 4),
"miss_rate": round(miss / routed, 4),
# A预测到却没进池budget/时序/驱逐B没预测到召回缺口
"miss_A_predicted": M["miss_A_predicted"],
"miss_B_unpredicted": M["miss_B_unpredicted"],
"A_share_of_miss": round(M["miss_A_predicted"] / max(1, miss), 4),
"B_share_of_miss": round(M["miss_B_unpredicted"] / max(1, miss), 4),
}
if M["dec_n"]:
# decode/verify 热路径专用(与上面 prefill-主导的全量分开看):
dr = max(1, M["dec_routed"])
dmiss = M["dec_miss_A"] + M["dec_miss_B"]
out["decode"] = {
"routed": M["dec_routed"],
"hit_rate": round(M["dec_resident_hit"] / dr, 4),
"miss_rate": round(dmiss / dr, 4),
"miss_A_predicted": M["dec_miss_A"],
"miss_B_unpredicted": M["dec_miss_B"],
"A_share_of_miss": round(M["dec_miss_A"] / max(1, dmiss), 4),
"B_share_of_miss": round(M["dec_miss_B"] / max(1, dmiss), 4),
# miss_A 细分:时序(pread 没完成) vs 驱逐(就绪过但 acquire 前不在池)
"miss_A_timing": M["dec_miss_A_timing"],
"miss_A_evicted": M["dec_miss_A_evicted"],
"A_timing_share": round(M["dec_miss_A_timing"] / max(1, M["dec_miss_A"]), 4),
"A_evicted_share": round(M["dec_miss_A_evicted"] / max(1, M["dec_miss_A"]), 4),
}
return out
def _predict_recall():
from mlx_streaming.core.profiling import PREDICT_RECALL_PROF as P
if not P["n"]:
return None
return {"recall": round(P["hit"] / max(1, P["routed"]), 4),
"avg_routed": round(P["routed"] / P["n"], 2), "n": P["n"]}
def _window_prof():
from mlx_streaming.core.profiling import WINDOW_PROF
if not WINDOW_PROF["n"]:
return None
return {"avg_ms": round(WINDOW_PROF["sum_s"] / WINDOW_PROF["n"] * 1000, 4),
"n": WINDOW_PROF["n"]}
def _prefetch_tprof(wall_s=None):
"""预取 host 墙钟探针汇总(PREFETCH_TPROF=1 才有数据)。
汇报每段:总秒数、占最终测量轮 wall 的百分比、单次调用均值(ms)。
注意:这是"主线程不可重叠的 host 时间";gate matmul / pool scatter 的 GPU 执行落在
前向末尾统一 eval,不在此口径,需用消融(NATIVE_NO_SUBMIT/NATIVE_NO_PROMOTE)差值测。
"""
from mlx_streaming.core.profiling import PREFETCH_TPROF as T
if not (T["predict_n"] or T["promote_n"] or T["submit_n"]):
return None
def _seg(s, n):
d = {"total_s": round(T[s], 4)}
if wall_s:
d["pct_wall"] = round(T[s] / wall_s * 100, 2)
if n:
d["avg_ms"] = round(T[s] / n * 1000, 4)
return d
host_total = T["predict_s"] + T["submit_s"] + T["promote_s"]
out = {
"predict": _seg("predict_s", T["predict_n"]),
"submit": _seg("submit_s", T["submit_n"]),
"promote": _seg("promote_s", T["promote_n"]),
# promote 内部细分(take 锁读 / route 成员同步 / place 切片+scatter 入图)
"promote_take": _seg("take_s", T["promote_n"]),
"promote_route": _seg("route_s", T["promote_n"]),
"promote_place": _seg("place_s", T["promote_n"]),
"host_total_s": round(host_total, 4),
"host_total_pct_wall": (round(host_total / wall_s * 100, 2) if wall_s else None),
"counts": {"predict_n": T["predict_n"], "submit_n": T["submit_n"],
"promote_n": T["promote_n"], "place_experts": T["place_experts"]},
}
return out
if __name__ == "__main__":
main()