Prox Deep Dive Blog
Prox 完全拆解:先用便宜代理找通道,再用精确计算守住 SwiGLU 的质量 先问一个部署问题:70% 的 FFN 通道可以跳过,但谁来决定该跳过谁? 在小 batch、自回归解码中,Transformer 的瓶颈经常不是注意力,而是把权重从 HBM 搬到片上存储。现代 LLM 普遍使用 SwiGLU FFN:up、gate、down 三个矩阵共同占据了 FFN 的参数、内存流量和乘加运算。以 Qwen3-8B 为例,$d_{\rm model}=4096$、$d_{\rm ff}=12288$,每层 FFN 有 $$ 3d_{\rm model}d_{\rm ff} =3\times4096\times12288 \approx 1.51\times10^8 $$个权重参数。逐 token 解码时,这三次投影都要反复读取,因此激活稀疏是一个高杠杆的优化点。 但“把 70% 的值置零”并不自动等于“质量只损失 70%”。每个中间通道还会经过 down 投影的不同权重行;错误的通道选择可能比同样数量的正确选择昂贵得多。更棘手的是,真正有用的信号恰好是 SwiGLU 的中间状态,而它只有在执行 up 和 gate 之后才知道:如果先做完整计算再选择通道,就已经失去了大部分加速机会。 这篇稿件(Jinyi Liu 等,来源文档未给出正式 venue 和发布日期)围绕一个中心问题展开:怎样以远低于 dense FFN 的代价,提前得到足够可靠的中间通道掩码? 一句话概括 Prox:用输入稀疏加 INT4 权重构造一个只负责“排名”的代理状态,再用原始权重对入选通道做精确计算。 SwiGLU 中间状态为什么是最自然的选择信号? 给定输入 $\mathbf{x}\in\mathbb{R}^{d_{\rm model}}$,SwiGLU 的三个中间量为 $$ \mathbf{u}=\mathbf{x}W_{\rm up},\qquad \mathbf{h}=\operatorname{SiLU}(\mathbf{x}W_{\rm gate}),\qquad \mathbf{s}=\mathbf{u}\odot\mathbf{h}, $$最终输出是 $\mathbf{y}=\mathbf{s}W_{\rm down}$。把它改写成按通道求和: ...