arXiv'26 | APEX:端侧 MoE 的专家预取,别再 top-k 一刀切了

arXiv’26 | APEX:端侧 MoE 的专家预取,别再 top-k 一刀切了

原文:APEX: Adaptive Expert Prefetching for Edge MoE Inference


1. 前言:先把背景交代清楚

你有没有想过这样一个问题:MoE 模型号称”稀疏激活、天生适合端侧”,为什么真放到手机/边缘盒子上,还是慢得像老牛拉车?

先说清楚 MoE 的工作方式,这是理解全文的前提。每个 token 进到 MoE 层,router(一个很小的分类头)从 N 个专家里挑出 top-k 个(比如 Granite-3B 是 40 选 8),只有这 k 个专家的 FFN 参与计算——这就是”稀疏激活”:计算量只跟 k 走,跟 N 无关。所以一个 3B 总参的 MoE,每 token 实际只动 800M 参数。

听起来是端侧的完美配方?坏就坏在“激活稀疏”不等于”加载稀疏”。端侧芯片的片上内存(SRAM/封装内存)只有几个 GB,装不下全部专家权重,绝大多数专家只能躺在片外的 LPDDR5X 里。于是每来一个 token,流程是:

route(挑专家)→ load(从 LPDDR5X 搬权重)→ execute(算)
                    ↑
              这一段是纯等待,计算阵列全程闲置

论文在 Granite-3.1-3B-A800M(1024 token 上下文)上量了这笔账:专家加载占延迟的 43%、总能耗的 29%。搬权重的时间比算权重的时间还长,而且计算阵列闲置期间还在烧静态功耗——花钱买了个电暖气。

端侧 MoE 的延迟与能耗拆解

标准解法是预取(prefetching):趁当前层还在算,提前把下一层要用的专家搬进来,把 load 藏进 execute 里。这条线我之前写过一篇:《ST-MoE expert-prefetch》——它的核心洞察是 expert 激活不是随机的,有时空相关性(相邻 token、相邻层倾向复用同一批专家),所以可以拿历史激活模式预测下一层要哪些专家,提前搬上片。另一篇 《SMoE》走的是完全不同的路:专家不在 GPU 上就别死等,找个”平替”专家直接算——用一点精度换完全不 stall。

今天这篇 APEX(arXiv 2608.11688)要指出的是:现有预取的”固定 top-k”策略,在端侧是一个两头堵的死局,而它的答案是让预取预算本身随 token 自适应。在模拟的高端端侧平台上:重叠准确率保住 97–98%(ProMoE 会掉到 79%),延迟比无预取低 42%,能耗比固定超采低 21.8%,EDP 最高改善 41%。


2. 病根:平均准确率是个谎言

固定 top-k 预取的代表是 ProMoE:预取时多搬 δ 个”预测最可能的”专家。它的问题是 δ 是个全局常数

  • δ 给小了(比如 +2):难 token 的预测 miss,router 要的专家不在片上,照样 stall;
  • δ 给大了(比如 +8):简单 token 白搬一堆用不上的专家,能耗爆炸。

你可能会问:预测准确率 70–85%(ProMoE 的水平)听起来也不差啊?这里有个很隐蔽的度量陷阱。论文强调 overlap 必须按 per-layer、per-token 算:每个 token 的每一层都要求”router 选的 k 个专家 ⊆ 预取集合”才算成功。为什么这么苛刻?因为平均数会骗人——k=4 时哪怕每层平均重叠 75%(4 个里对 3 个),每个 token 每层都缺一个专家,每个 token 都 stall。串行系统里,99 层完美 + 1 层 miss = 全部白干,这一点我在《KVCOMM》那篇讲依赖链的时候就提过:流水线的吞吐由最慢的一环决定,不由平均环决定

而 token 之间的”预测难度”差异极大——论文 profiling 发现:大多数 token 只要 δ=2 或 δ=4 的小额预算就能全覆盖,少数难 token 才需要大预算。这个长尾分布就是”固定 δ 必然两头堵”的实证根据,也是自适应方案的立足点。


3. 平台:先把”端侧设备”说清楚

这是篇算法-硬件协同设计的论文,评估平台是 CHIPSIM 综合校准联合仿真器(不是真机)。读这类论文第一步永远是搞清楚平台假设,这里的目标设备定位在”近端旗舰边缘加速器”(Jetson Thor / Hailo-10H 一档),chiplet 架构:

计算 chiplet

  • 4 个 vector processing array,每个是 16×16 的 vector unit 阵列、vector length 32,BF16 matmul @ 750MHz → 峰值 24 TFLOPS
  • 每阵列 2MB SRAM(共 8MB)+ 32 个 SFU 处理非线性;
  • 阵列是 RTL 级实现、TSMC 28nm 综合、Synopsys PrimeTime 功耗表征——不是拍脑袋的性能模型。

内存层级(这是全文的关键变量):

  • 封装内:HBM3-class 819 GB/s(7 pJ/bit)——放共享的热数据;
  • 片外:专家权重在 LPDDR5X(3 pJ/bit),走 PCIe 6.0 ×16、256 GB/s(5 pJ/bit);
  • 片间 UCIe 0.5 pJ/bit。

预取按真实 DMA 建模——受带宽竞争和排队延迟影响,不是理想化的”并行搬运不要钱”。内存用 Ramulator 2.0 建模,整体是 cycle-accurate 的联合仿真。这个严谨度比”在 A100 上模拟手机”高一档。


4. 方法:一个会看人下菜的预取路由器

APEX 的 pipeline 把 route → load → execute 改成 predict → prefetch → execute,由三个组件构成。

4.1 Prefetch Router:预测本层专家

一个线性层 + softmax,从冻结的原 router 蒸馏(KL loss)。有两个设计决策值得划重点:

(1)它放在本层 attention 之前,用当前层的 hidden state 预测本层 FFN 将要路由的专家。这和 Pre-gated MoE 的”用第 L 层预测第 L+1 层”不同——APEX 的理由是同层内的 hidden state 和最终路由决策相关性最强(跨层预测等于又叠了一层不确定性)。而且预取发起得早,DMA 传输可以藏进 attention 的计算时间里。

(2)原 router 保留,依然是最终决策者。prefetch router 只负责内存调度(搬谁),不参与计算路由(用谁)——这样预测错了最多是多搬/漏搬,不会改变模型的计算语义。这一点和 Pre-gated MoE 直接用预测结果做路由的激进路线形成对比。

蒸馏 loss 就是标准 KL:

\[\mathcal{L}_\text{KL} = \sum_i q_r(i) \log \frac{q_r(i)}{q_p(i)}\]

$q_r$ 是原 router 的专家分布,$q_p$ 是 prefetch router 的,梯度只更新后者,基座模型完全不动

4.2 自适应预算 δ̂(x):这篇的灵魂

固定 δ 的替代方案:给每个 token 实时决定”多取几个”。先定义 oracle——δ*(这个 token 的真实难度):

\[\delta^* = \min\{\delta \in \{0, \dots, N-k\} : \mathcal{K}_r \subseteq \mathcal{K}_p^{(\delta)}\}\]

即”让 router 真正选的 k 个专家全部落进 top-(k+δ) 预取列表”所需的最小额外预算。δ*=0 表示预测全对,δ* 大表示这个 token 难猜。

然后训一个序数 logistic CDF 模型来预测 δ* 的分布:

\[p_\delta(x) = \Pr(\delta \geq \delta^* \mid x) = \sigma(\theta_\delta - w^\top x)\]

$x$ 是该层 attention 之前的 token 表征,$w$ 是一个权重向量(就这么简单,一个线性打分器),$\theta_0 \leq \theta_1 \leq \dots \leq \theta_{N-k}$ 是有序阈值。推理时的规则:

\[\hat{\delta}(x) = \min\{\delta : p_\delta(x) \geq \tau\}\]

“取最小的 δ 使覆盖率有 τ 的把握”——τ=0.90 是默认值。有把握的 token 只多取一两个,没把握的 token 才大手笔超采,预算精确花在长尾上。

训练分两阶段(都是离线的,且训一次全任务通用):

# Stage 1: KL 蒸馏 prefetch router(见上式)
# Stage 2: 用累积二元交叉熵拟合 CDF
def cdf_loss(theta, w, x, delta_star):
    # delta_star: profiling 得到的 oracle 标签
    # 对每个候选预算 delta,目标是 1[delta >= delta_star]
    loss = 0
    for delta in range(0, N - k + 1):
        p = torch.sigmoid(theta[delta] - w @ x)        # P(delta* <= delta)
        target = 1.0 if delta >= delta_star else 0.0
        loss += F.binary_cross_entropy(p, torch.tensor(target))
    return loss

训练配置很轻:WikiText 上 lr 5e-4、batch 8、序列长 1024、1000 步,每个模型 10 分钟到 1 小时。整个 APEX 新增参数 0.79M–34.11M,只占模型参数的 0.022%–0.060%,推理性能开销 <0.051%。

4.3 两种执行模式:语义严格 vs 绝不等待

预取 miss 了怎么办?两种模式:

Correctness-preserving(默认):router 要的专家有缺失时,先把已就绪的专家算起来,缺失的异步补取,聚合前补齐——最终聚合用的是完整的原始专家集合,语义与原 MoE 严格一致,代价只是没藏住的纠错延迟。

Stall-free(可选):绝不等待。缺失的专家直接用”已预取集合里、原 router softmax 分数最高的候选“顶替——这就是 SMoE 那篇”平替”思想的端侧版。Algorithm 1 的逻辑:

def stall_free_select(K_routed, q_r, K_prefetched):
    # K_routed: 原 router 选的 k 个专家; q_r: 原 router 的 logits
    # K_prefetched: 按 delta_hat(x) 预取的 k+delta 个专家
    alpha = softmax(q_r)                          # 用原 router 的权重当替补排序依据
    if K_routed <= K_prefetched:                  # 全命中,直接用
        return K_routed
    keep  = K_routed & K_prefetched               # 预测对的保留
    missing = K_routed - K_prefetched             # 缺失的
    subs  = top_n_by_weight(K_prefetched - keep, alpha, n=len(missing))
    return keep | subs                            # 用预取池里分数最高的顶上

注意替补的排序依据是原 router 的分数(不是 prefetch router 的)——原 router 才知道”哪些专家真的合适”,prefetch router 只管”搬谁”。


5. 效果

5.1 重叠准确率:自适应的碾压局

τ=0.90 下,APEX 在所有层保持 >97–98% 的专家重叠率;ProMoE 在 Granite-1B 的第 5 层会掉到 79%。平均每个 token 额外预取的专家数:Granite-1B 4.17、Granite-3B 2.86、Phi-7B 0.67、DeepSeek-V2-Lite 1.98——模型越”好预测”,预算自动越小,这就是自适应的意义(Phi 只有 16 个专家、路由集中,0.67 个额外预算就够 98%)。

τ 的扫描直接展示”把握-预算”的交换关系:Granite-1B 上 τ 从 0.60→0.90→0.97,重叠率 93.4%→98.2%→99.4%,平均 δ̂ 从 1.36→4.17→6.70——想要多稳,明码标价

重叠准确率随预算的变化

5.2 延迟与能耗:预取生意的本质

Granite-3B、512 token、correctness-preserving 模式:

  • 延迟 11.41ms:比无预取(19.77ms)低 42%,比 ProMoE(15.39ms)低 26%,比固定超采 (k+4)/(k+8) 分别低 20%/40%;
  • 能耗 287.3mJ:比无预取低 9.5%、比 ProMoE 低 5.8%;固定 k+8 超采最惨(比 APEX 多 21.8%)——固定预算省下的延迟全在能耗上还回去了;
  • Stall-free 再加 2.0–2.8% 的延迟改善,聊胜于无。

延迟与能耗对比

能耗的拆解(1024 context)最能说明”预取生意”的本质:开 APEX 后,片外 I/O 从 48mJ 涨到 56mJ(多搬了 8mJ),但计算阵列的闲置漏电从 57mJ 崩到 8mJ(少烧 49mJ)——用 8mJ 的搬运换 49mJ 的干等,净赚 41mJ。这就是所有预取技术的账本:带宽是便宜的,闲置是昂贵的。

5.3 鲁棒性与精度代价

  • EDP(延迟×能量积):相对无预取改善 36–49%、相对 ProMoE 改善 16–30%(视模型而定),Phi-7B 上 stall-free 的增益尤其大(再 +14/8/7%)——因为它 k=2,miss 一个就废一个,绝不等候的价值最大;
  • 带宽扫描:32–1024 GB/s 全区间延迟收益 14–42% 成立——带宽越紧张,预取越值钱;
  • 精度扫描:4/8/16/32-bit 专家权重下收益都在——和量化正交可叠加(量化缩体积、APEX 缩次数);
  • Stall-free 的精度代价(明码标价):Granite-1B PPL 7.88→7.99、平均分 43.3→42.8;DeepSeek-V2-Lite PPL 7.02→7.10。Phi-7B 掉得明显(64.6→60.4)——16 选 2 的模型每个专家都是命根子,替代一个就伤筋动骨。论文因此建议 k 小的模型关掉 stall-free,用 correctness-preserving——这个诚实的边界标注很加分。

6. 我的 Take

一,预取问题的正确度量是尾部覆盖率,不是平均准确率。 ProMoE 的 80% 平均准确率听起来能打,但任何一层的 miss 都会把前面所有层的完美预取全部白干。”平均”在串行依赖链上是个没有意义的统计量——这个认识对做流水线/调度系统的同学是通用的。APEX 用一个 ordinal logistic CDF 把”这个 token 需要多少预算”建模出来,本质上是把置信度校准这套工具用到了内存搬运调度上,这个迁移本身就很漂亮。

二,端侧 MoE 的解题空间基本铺开了,三个角各站一篇。 预测得准(ST-MoE 时空相关性)、取得巧(APEX 自适应预算)、算得动(SMoE 平替)——三者还是正交的:APEX 的 prefetch router 可以吃任何更好的预测器,stall-free 的替补池也可以换成 SMoE 式的更聪明的替代选择。往后端侧 MoE 的论文大概率是在这三个角之间做组合。

三,对”仿真论文”的一点辩护。 这篇没有真机,但 RTL 综合 + PrimeTime 功耗 + Ramulator 内存 + DMA 排队的建模链,比多数”在服务器 GPU 上假装端侧”的论文可信得多。硬件协同设计类工作本来就该先在仿真里把设计空间扫清楚——当然,上真机验证仍是它(和所有同类工作)欠社区的一步。


如果这篇文章涉及的 MoE 推理优化你想系统深入,可以看看我之前出版的《动手学AutoML:从 NAS 到大语言模型优化实战》,书里讲 NAS 在大模型上的延伸时正好覆盖稀疏模型结构搜索这条线,和 MoE 效率优化是同一套设计思路。

动手学AutoML书籍封面

Flag Counter