混合线性注意力
混合线性注意力大模型(HLA LLM)介绍
背景与动机
Transformer 架构中的标准自注意力(Full Attention)机制虽然效果强大,但计算复杂度随序列长度呈平方级增长(O(n²)),这在处理长序列时带来严重的计算和显存开销。线性注意力(Linear Attention) 通过核函数近似、状态压缩等方式将复杂度降至 O(n),大幅提升效率,但通常在长距离依赖建模、精确检索(如"大海捞针"任务)等能力上弱于全注意力。
混合线性注意力(Hybrid Linear Attention, HLA) 类模型的核心思路,就是在同一个网络中结合两种机制,试图"取其精华":用线性注意力承担大部分计算以提效,同时保留部分全注意力模块或机制来补足能力短板。
核心设计思路
1. 层间混合(Layer-wise Hybrid)
最常见的方式,在网络的不同层间交替或按比例分配两种注意力:
- 大部分层使用线性注意力(如 Mamba、RWKV、GLA、RetNet 等的核心机制)
- 少数关键层(通常是固定间隔,如每 N 层)保留标准全注意力
- 代表工作:Jamba、Zamba、Griffin、Samba 等都采用了类似思路
2. 头内混合(Head-wise Hybrid)
在同一层的多头注意力中,部分头做线性计算、部分头做全注意力,让每层都能兼顾局部效率与全局能力。
3. 状态增强型线性注意力
在线性注意力的循环状态(recurrent state)基础上,引入额外机制增强其记忆容量和检索能力,例如:
- 增加状态维度或多状态并行(如 DeltaNet 的 delta rule 更新)
- 引入门控机制动态调节信息保留与遗忘
- 结合滑动窗口注意力(Sliding Window Attention)作为局部全注意力补充
4. 训练后蒸馏/转换
将已训练好的全注意力大模型通过参数映射或蒸馏,转换出一个混合结构模型,减少从头训练成本(如 Mamba-in-Llama、Llamba 等尝试)。
兼顾效率与能力的关键机制
| 机制 | 作用 |
|---|---|
| 固定比例混合层 | 控制计算量的同时保留全局建模能力 |
| 滑动窗口 + 线性注意力 | 局部精确 + 全局压缩,性价比高 |
| 门控/衰减机制 | 提升线性注意力对关键信息的选择性保留 |
| KV Cache 压缩 | 全注意力层的显存开销也被针对性优化 |
| 推理时的状态复用 | 线性部分保持 O(1) 的逐 token 推理开销 |
优势
- 推理效率提升:相比纯全注意力模型,显存占用和长序列推理速度显著改善,尤其在超长上下文场景(如百万 token 级别)
- 能力损失可控:通过保留少量全注意力层,在"大海捞针"、多步推理、长距离依赖等任务上的表现接近纯全注意力模型
- 训练/部署灵活性:可以根据硬件和场景需求调整线性层与全注意力层的比例,形成效率-能力的可调谱系
挑战与局限
- 架构设计的经验性较强:混合比例、放置位置目前多依赖实验搜索,缺乏系统性理论指导
- 长距离精确检索仍有差距:纯线性注意力层在处理精确 token 级检索任务上仍弱于全注意力,混合模型的上限受限于全注意力层的比例和位置
- 训练稳定性:两种机制的梯度特性、数值尺度不同,混合训练可能需要额外的调参和归一化设计
- 生态成熟度:相比标准 Transformer,配套的推理优化(如 vLLM、TensorRT 支持)和调试工具链还在发展中
典型代表模型(架构思路参考)
- Jamba(AI21):Transformer + Mamba 混合,比例约为 1:7
- Zamba:Mamba 主干 + 共享注意力层
- Griffin / RecurrentGemma(Google DeepMind):线性递归 + 局部注意力
- MiniMax-01系列:闪电注意力(Lightning Attention,线性注意力变体)与全注意力混合
- Qwen3-Next等新一代模型也在探索类似的混合架构方向
总结
HLA 类模型代表了大模型架构设计中"效率与能力权衡"的一个重要方向:不再追求单一注意力机制的极致,而是通过结构化混合,在超长上下文处理、推理成本和模型能力之间寻找更优的帕累托前沿。这也反映了当前大模型研究从"堆参数堆算力"向"架构效率创新"转变的趋势。
如果你对某个具体模型(如 Jamba 或 MiniMax)的架构细节感兴趣,我可以进一步展开介绍。
线性注意力(Linear Attention)详解
一、为什么需要线性注意力
标准自注意力的计算过程:
问题在于 \(QK^T\) 这一步:\(Q, K \in \mathbb{R}^{n \times d}\) 相乘得到 \(n \times n\) 的矩阵,计算复杂度 O(n²d),显存占用 O(n²)。当序列长度 n 达到几万甚至几百万时,这个开销是致命的。
线性注意力的目标:去掉 softmax 这个非线性瓶颈,利用矩阵乘法的结合律,把复杂度降到 O(nd²),即随序列长度线性增长。
二、核心数学技巧:去掉 Softmax,利用结合律
标准注意力的本质
可以把 softmax 注意力写成核函数(kernel)形式:
其中 \(\text{sim}(q_i,k_j) = \exp(q_i \cdot k_j / \sqrt{d})\)。
线性注意力的替换
将相似度函数替换为可分解的核函数:
这里 \(\phi(\cdot)\) 是某个特征映射函数(可以是恒等映射、elu+1、随机特征映射等)。
关键在于,一旦相似度可以分解,就能利用矩阵乘法的结合律重新排列计算顺序:
左边先算 \(QK^T\) 得到 \(n\times n\) 矩阵再乘 V,复杂度 O(n²d);
右边先算 \(\phi(K)^T V\) 得到 \(d \times d\) 矩阵,再左乘 \(\phi(Q)\),复杂度 O(nd²)。
这就是线性注意力效率提升的数学本质:把"先乘 Q,K 再乘 V"变成"先乘 K,V 再乘 Q"。
三、递归形式:与 RNN 的等价性
线性注意力有一个重要性质——可以写成递归(recurrent)形式,这也是它能在推理时做到 O(1) 逐 token 开销的关键。
定义状态矩阵:
则输出:
这本质上就是一个线性 RNN:状态 \(S\) 以 \(d \times d\)(或 \(d \times d_v\))的固定大小滚动更新,不随序列长度增长。这正是线性注意力能做到推理时显存恒定、速度恒定的原因,也是它和 Mamba、RWKV、RetNet 等"现代 RNN"能够统一理解的桥梁。
四、代表性方法演进
1. Linear Transformer(2020)
最早的工作之一,用 \(\phi(x) = \text{elu}(x) + 1\) 作为特征映射,验证了上述结合律技巧的可行性。
2. Performer
用随机特征(Random Features)近似 softmax 核,理论上更贴近原始 softmax,但实现复杂、方差较大。
3. RetNet(Retentive Network)
引入指数衰减的门控机制,让状态更新带有"遗忘"性质:
\(\gamma < 1\) 使得早期信息逐渐衰减,缓解线性注意力"记忆无限累积、无法区分远近"的问题,同时保留了并行计算的能力(可以用 chunk-wise 并行)。
4. GLA(Gated Linear Attention)
把衰减系数从标量 \(\gamma\) 扩展为逐维度、数据依赖的门控向量:
其中 \(\alpha_i\) 由输入动态生成,让模型自适应地控制每个维度的信息保留/遗忘速度,表达能力更强。
5. DeltaNet
不是简单累加信息,而是用Delta Rule(类似梯度下降/纠错)更新状态:
这可以理解为对状态做"覆盖修正"而非单纯"叠加",缓解了信息随时间被稀释、旧信息覆盖新信息不准确的问题,检索能力比纯累加式的线性注意力更强。
6. Mamba(State Space Model 视角)
严格来说是结构化状态空间模型(SSM),但与线性注意力高度同构:
通过选择性机制(Selective SSM)让 \(A, B, C\) 依赖于输入动态变化,是当前"线性化序列建模"家族中效果最突出的分支之一。
五、训练与推理的两种计算模式
线性注意力类模型的一大优势是同一模型可以用两种等价方式计算:
| 模式 | 适用场景 | 特点 |
|---|---|---|
| 并行模式(Chunk-wise/Parallel) | 训练阶段 | 利用矩阵乘法批量计算,充分利用 GPU 并行性,类似分块 attention |
| 递归模式(Recurrent) | 推理阶段(逐 token 生成) | 只维护一个固定大小的状态 \(S\),每步 O(1) 更新,无需 KV Cache 随长度线性增长 |
这种"训练并行、推理递归"的双模式设计,是 RetNet、GLA、Mamba 等模型的通用范式,兼顾了训练效率与推理效率。
六、线性注意力的局限性
- 表达能力上限受状态大小限制:\(d \times d\) 的状态矩阵是信息压缩的瓶颈,序列越长,早期信息越容易被"稀释"或覆盖,不像全注意力可以无损保留所有 KV 对
- 精确检索能力弱:在需要精确定位某个 token(如"大海捞针"任务)时,压缩状态难以保证不丢失关键信息,这也是 Full Attention 仍不可替代的原因
- 不同变体效果差异较大:核函数选择、门控设计、状态更新规则的细节对最终效果影响显著,尚无统一的"最优解"
- 硬件适配仍在发展:虽然理论复杂度更低,但由于矩阵形状(如 \(d\times d\) 的小矩阵操作)不如标准 attention 在 GPU 上"规整",实际加速比要靠专门 CUDA kernel(如 flash-linear-attention 库)才能充分发挥
七、小结:一条清晰的演进脉络
Linear Transformer(去 softmax,结合律加速)
↓
RetNet(引入标量衰减门控,缓解遗忘问题)
↓
GLA(衰减门控从标量升级为逐维度、数据依赖)
↓
DeltaNet(从"累加"升级为"纠错式覆盖",增强检索能力)
↓
Mamba/SSM(从注意力视角转向状态空间视角,选择性机制进一步增强)
这条脉络也解释了为什么 HLA(混合架构)会成为一个自然的折中方案:纯线性注意力在状态压缩上的固有局限,短期内难以通过单一机制彻底解决,因此工程实践上选择"用少量全注意力层弥补短板",而不是押注某一种线性注意力变体能完全替代全注意力。
Q、K、V 的维度详解
一、基础设定
假设一个标准的多头注意力层,涉及以下超参数:
| 符号 | 含义 | 常见取值示例 |
|---|---|---|
| \(b\) | batch size(批大小) | 32 |
| \(n\) | 序列长度(token 数) | 4096 |
| \(h\) | 注意力头数(head 数) | 32 |
| \(d\) | 每个头的维度(head dim) | 128 |
| \(d_{model}\) | 模型隐藏维度 | \(h \times d = 4096\) |
二、单头情况下的维度(方便理解)
先忽略 batch 和多头,只看单个样本、单个头的情况,序列长度为 \(n\),特征维度为 \(d\):
通常 \(d_v = d\)(K 和 V 的维度可以不同,但实践中大多设成相同)。
标准全注意力的维度变化
这一步产生了 \(n \times n\) 的注意力矩阵,这正是 O(n²) 复杂度的来源。
最终输出维度是 \(n \times d_v\),和输入 Q 的 token 数一致。
三、线性注意力中的维度变化(关键对比)
线性注意力通过特征映射 \(\phi\) 和结合律,把计算顺序换成先算 \(K^TV\):
这里维度变化的关键点:
- \(\phi(K)^T V\) 的结果不再是 \(n \times n\),而是一个固定大小 \(d \times d_v\) 的矩阵,与序列长度 \(n\) 无关!
- 这个 \(d \times d_v\) 矩阵,就是前面提到的状态矩阵 \(S\)
再用 \(\phi(Q)\) 左乘:
最终输出仍是 \(n \times d_v\),和全注意力结果维度一致,但中间从未出现过 \(n \times n\) 的矩阵,这就是省下 O(n²) 显存和计算的关键所在。
四、递归形式下,状态 S 的维度
在逐 token 处理(推理阶段)时:
这里单个 token 的维度:
- \(k_i \in \mathbb{R}^{d}\)(列向量)
- \(v_i \in \mathbb{R}^{d_v}\)(列向量)
- \(\phi(k_i) v_i^T \in \mathbb{R}^{d \times d_v}\)(外积,秩为1的矩阵)
- \(S_i \in \mathbb{R}^{d \times d_v}\)(与 \(S_{i-1}\) 维度相同,逐步累加更新)
这就是为什么线性注意力推理时显存恒定:无论生成到第几个 token,状态 \(S\) 始终是固定的 \(d \times d_v\) 大小,不会随 \(n\) 增长(对比全注意力的 KV Cache,会随 \(n\) 线性增长)。
五、加入多头和 batch 后的完整维度
实际实现中,张量形状通常是:
全注意力:
线性注意力(并行/chunk-wise 模式):
线性注意力(递归模式,推理时):
状态张量:
每一步更新时,输入单个 token 的 \(q_i, k_i, v_i \in \mathbb{R}^{b \times h \times d}\)(或 \(d_v\)),更新后 S 形状不变。
六、维度对比总结表
| 全注意力(Full Attention) | 线性注意力(Linear Attention) | |
|---|---|---|
| 中间矩阵 | \(n \times n\)(随序列长度平方增长) | \(d \times d_v\)(固定大小,与 n 无关) |
| KV Cache(推理) | 随 \(n\) 线性增长:\(O(n \cdot d)\) | 固定大小状态:\(O(d \times d_v)\) |
| 计算复杂度 | \(O(n^2 d)\) | \(O(n d^2)\) |
| 单步推理开销 | 需要和历史所有 KV 做计算 | 只需更新固定大小状态 |
七、直观理解:为什么维度设计带来效率提升
可以这样理解两种方式的本质差异:
- 全注意力:把"序列长度"这个维度放在了中间结果的两个轴上(\(n \times n\)),序列越长,中间结果越大
- 线性注意力:把"特征维度 d"放在了中间结果的两个轴上(\(d \times d_v\)),这个大小由模型设计决定(通常固定为 64、128 等),与序列长度完全解耦
这也是为什么线性注意力能做到"序列多长都不怕",但代价是 \(d \times d_v\) 这个"记忆容量"是有限的、固定的——这正是第六部分提到的"状态压缩瓶颈"的维度本质。

浙公网安备 33010602011771号