注意力替代方案与混合专家
对应课程 Lecture 4:Attention alternatives and mixture of experts(Tatsu)
标准 Transformer 有两大“原罪”:
- 注意力的
复杂度:序列翻倍,计算量翻四倍,长上下文贵得离谱; - 稠密激活:每个 token 都要经过全部 FFN 参数,推理时每生成一个词都要“动用”整个模型。
这一讲讨论这两个问题的主流解法。
注意力的成本分析
回顾注意力:token
思考一个尖锐的问题:每个 token 真的需要“看见”所有历史 token 吗?
- 翻译一个句子,大部分注意力集中在局部;
- 很多头学到的模式非常简单(关注最近的几个词、关注句首);
- 检索类任务(大海捞针)确实需要全局访问,但那是少数头少数层的事。
这就打开了替代方案的设计空间。
注意力的替代方案
局部注意力:滑动窗口
滑动窗口注意力(sliding window attention) 只让每个 token 关注最近
- 计算复杂度
,KV 缓存 ,都是线性; 取几千,直觉上“几千 token 前的东西对当前词影响很小”; - 叠多层后感受野线性扩大(每层 +
),深堆叠仍能看到远方。
代价:远距离信息必须一层层传,纯局部模型在长程检索任务上明显吃亏。
线性注意力与状态空间模型
另一条路线干脆取消“每对 token 两两计算”,用一个固定大小的状态(state) 递推地压缩历史:
这就是线性注意力/循环模型一族的通用形态,状态空间模型(State Space Model,SSM)(如 Mamba)是其中的代表:
- 训练可以并行(类似卷积),推理像 RNN 一样
每步; - 推理时状态恒定:不管上下文多长,每步计算和显存不变,吞吐极高;
- 代价:有限大小的状态装不下无限历史,精确检索能力弱于全注意力。
混合架构:取长补短
实践中的胜出方案往往不是“替代”而是“混搭”:把少数几层全局注意力插进大量 SSM/滑窗层中,比如 1:5~1:7 的比例。代表工作有 Jamba、Zamba,以及逐层混用的 Command 系列。
怎么理解混合架构的性价比?
全注意力保证“精确检索”能力的下限;便宜层负责大头的信息加工。用 10% 的注意力成本买到接近全注意力的质量,是典型的帕累托改进。选层时全局注意力通常放在靠后和靠前的位置(开头负责建立编码,结尾负责汇总输出)。
稀疏注意力
保留注意力形式,但让每个 token 只关注选出来的子集:
- 块状稀疏:只算对角块 + 若干全局列(如 Longformer 的模式);
- 硬件对齐分块:NSA、MoBA 等按块选择 KV,兼顾稀疏与 GPU 利用率;
- 难点:稀疏模式 irregular 时,GPU 利用率掉得厉害,省的理论算力可能被访存吃掉。
混合专家:稀疏激活参数
核心思想
混合专家(Mixture of Experts,MoE) 攻击的是第二个原罪:参数多 ≠ 每 token 都要用全部参数。
把每层的单个 FFN 换成
效果:总参数量可以堆到几百 B,但每个 token 只激活其中一小部分——参数容量大,计算量小。DeepSeek-V3、Mixtral 都是这一路线。
负载均衡:MoE 的头号难题
路由是“赢者通吃”的:如果几个专家特别好用,大家全往那儿挤,其他专家得不到训练,系统退化为小稠密模型。解决办法是加辅助负载均衡损失(auxiliary load balancing loss),惩罚路由分布的不均匀:
其中
其他关键工程细节:
- 容量因子(capacity factor):每个专家的缓冲区按“平均负载 × 容量因子”预分配,超出就丢弃(token dropping);
- 专家并行(expert parallelism):专家分布在不同 GPU 上,路由需要跨卡通信(all-to-all),这是并行策略的伏笔;
- 路由稳定性:fp32 精度算 softmax、_router_bias、辅助无损均衡等技巧都用于防止路由崩塌。
MoE vs 稠密:怎么选?
| 维度 | 稠密模型 | MoE 模型 |
|---|---|---|
| 同等训练算力的质量 | 基线 | 通常更好(参数更多) |
| 推理显存 | 小 | 大(全部专家都要驻留) |
| 推理带宽/延迟 | 每步激活全部参数 | 每步只激活部分参数 |
| 训练稳定性 | 成熟 | 路由不均衡、发散风险 |
直觉:MoE 拿“显存”换“质量/带宽”。在显存便宜的训练集群上划算,在显存受限的部署场景要掂量。
小结
- 注意力的
和 KV 缓存是长上下文的核心瓶颈;滑窗、SSM/线性模型、稀疏注意力从不同角度缓解。 - 实践赢家是混合架构:少数全局层 + 多数高效层。
- MoE 用稀疏激活把“参数容量”和“计算量”解耦;负载均衡损失和容量因子是训练成败的关键;代价是推理显存和工程复杂度。