Informer 长序列预测:用概率稀疏注意力砍复杂度
标准 Transformer 注意力是 O(L²):序列拉到 4096 步,一层注意力就要算 1600 万对相似度。Informer(AAAI 2021 最佳论文)的 ProbSparse 注意力抓住一个经验事实——注意力矩阵天然长尾,绝大多数查询行近乎均匀分布、对输出没贡献。它用采样估计每个查询的稀疏度 M(q,K)=max−mean,只让 top-u=c·lnL 个「活跃查询」做完整注意力,懒查询直接输出 V 的均值,复杂度砍到 O(L·lnL)。纯 numpy 从零实现 ProbSparse + 单层注意力回归(手写反向传播、有限差分梯度校验 1e-10 级),实测 L=4096 时提速 53 倍;但诚实披露:在 L=96 的短窗口预测任务上,注意力行不够长尾,c=10 的 ProbSparse 推理 R²=−0.40 远差于 full 推理 0.746,c 扫到 20(u/L≈96%)才恢复 0.727——稀疏近似的前提是分布真的稀疏。另附 OLS-96滞后 R²=0.834 反超注意力的诚实对照,拆穿「ProbSparse 无损/长序列必用 Transformer/加速免费/M 度量万能/金融直接落地」五类真实陷阱(中阶)。
你想用 Transformer 做长序列预测——比如喂 4096 根 K 线预测未来走势。立刻撞墙:标准自注意力要算每个查询对每个键的相似度,复杂度 O(L²)。L=4096 就是 1600 万对点积,一层就吃满显存,再深就训不动。
Informer(Zhou et al., AAAI 2021 最佳论文)给出的解法叫 ProbSparse 注意力,出发点是一个经验观察:注意力矩阵天然长尾。把训练好的注意力矩阵拿出来看,绝大多数查询行的分布几乎是均匀的——它对所有键”雨露均沾”,输出就约等于 V 的均值,算不算完整注意力根本无所谓。真正携带信息的,是少数分布尖锐的”活跃查询”。
既然如此:只给活跃查询算完整注意力,懒查询直接给均值,复杂度从 O(L²) 砍到 O(L·lnL)。

左图是我们合成结构化输入后的完整注意力矩阵:大片近均匀的行里嵌着少数亮条纹。右图把每个查询的稀疏度度量排序——典型的长尾:前 9% 的查询贡献了几乎全部”指向性”。
一、ProbSparse 的三步机制#
第 1 步:给每个查询打”稀疏度分”。 理想度量是查询行的注意力分布与均匀分布的 KL 散度,但算 KL 本身就要 O(L²)。Informer 用一个便宜的代理:
max 减 mean:分布越尖锐,最大值越突出于均值,M 越大。均匀分布的 M≈0。
第 2 步:采样估计 M。 精确算 M 还是要全部点积。关键技巧:每个查询只随机采 c·lnL 个键来估计 max 和 mean。理论依据是长尾分布下 max 大概率被少量样本命中(论文给了引理保证)。
第 3 步:top-u 完整注意力 + 懒查询均值。 取 M 最大的 u = c·lnL 个查询做完整 softmax 注意力,其余 L−u 个查询输出 V 的均值(自注意力里对应”均匀注意力的期望输出”)。
总代价:采样打分 O(L·lnL) + 头部注意力 O(L·lnL),整体 O(L·lnL)。
def probsparse_attention(Q, K, V, c=5, rng=None):
L, d = Q.shape
u = min(L, int(np.ceil(c * np.log(L)))) # 活跃查询数
n_sample = min(L, int(np.ceil(c * np.log(L)))) # 每查询采样键数
idx = rng.choice(L, size=n_sample, replace=False)
S_sample = Q @ K[idx].T / np.sqrt(d) # (L, n_sample)
M = S_sample.max(axis=1) - S_sample.mean(axis=1) # 稀疏度:max - mean
top = np.argsort(M)[-u:]
out = np.repeat(V.mean(axis=0, keepdims=True), L, axis=0) # 懒查询 → V均值
S_top = Q[top] @ K.T / np.sqrt(d) # 仅 u 行做完整注意力
out[top] = softmax(S_top, axis=-1) @ V
return out, M, toppython在 L=336 的演示里:活跃查询只占 8.9%,这些行的输出与完整注意力误差为 0(它们本来就是精确计算的);误差全部集中在懒查询上(行误差均值 0.70)。M 度量与真实”行熵偏离均匀”的相关性 0.548——便宜的代理,方向正确但并不完美。
二、加速是真的:L=4096 提速 53 倍#
同一份 numpy 代码实测前向耗时:
| L | Full (ms) | ProbSparse (ms) | 加速 |
|---|---|---|---|
| 512 | 1.41 | 0.21 | 6.8× |
| 1024 | 6.30 | 0.43 | 14.7× |
| 2048 | 25.2 | 0.78 | 32.3× |
| 4096 | 89.8 | 1.69 | 53.0× |

log-log 图上两条线斜率肉眼可分:full 是斜率 2 的抛物线,ProbSparse 近乎线性。L 越长,砍复杂度的收益越大——这正是 Informer 面向”长序列”的设计初衷。
三、预测实战:手写注意力回归 + 诚实的翻车现场#
合成 3600 步序列:三重周期(24/96/168 步)× 幅度调制 + 慢趋势 + AR(1) 噪声。窗口 L=96 预测下一步。模型是单层自注意力编码器(d=16,正弦位置编码,mean-pooling 读出),前向反向全部手写,有限差分梯度校验最大相对误差 6.7e-10——反向传播是对的。
结果(测试集 R²):
| 模型 | R² |
|---|---|
| 注意力(full 推理) | 0.746 |
| 注意力(ProbSparse 推理, c=10, u/L≈48%) | −0.400 |
| ProbSparse c=20(u/L≈96%) | 0.727 |
| OLS 96 滞后 | 0.834 |
| naive(上一值) | 0.651 |

两个扎心事实,都值得说透:
第一,ProbSparse 在这个任务上翻车了。 c=10 时只保留 48% 的活跃查询,R² 直接崩到负数;要 c=20(u/L≈96%,几乎不稀疏)才恢复到 0.727。为什么和第一节的”长尾”故事矛盾?因为这里的注意力矩阵不够长尾:L=96 的短窗口、平滑的周期信号,训练出的注意力行大多”温和地偏离均匀”——每行都有中等的信息量,没有谁可以被均值粗暴替代。ProbSparse 的前提是”大多数行真的近均匀”,这在 Informer 的目标场景(L 上千、编码器端)里常成立,在短窗口里不成立。稀疏近似的收益永远以分布真的稀疏为前提。
第二,OLS 赢了注意力。 96 个滞后的线性回归 R²=0.834 > 注意力 0.746。原因和我们在 TCN、因果卷积两篇里看到的一致:合成任务的主体是已知频率的线性叠加,线性模型在这类任务上接近贝叶斯最优。注意力的优势场景是变长依赖、内容寻址(“找相似历史片段”),而不是固定频率的周期外推。
四、五个真实陷阱#
陷阱一:「ProbSparse 是无损加速」。 不是。懒查询输出被均值替代,是有偏近似。整体输出相对误差在我们的结构化演示里是 49%(集中在懒行)。它赌的是”懒行的下游贡献本来就小”,这个赌注在注意力真长尾时才赢。
陷阱二:「长序列预测必须上 Transformer」。 我们的实验里 OLS 干掉了注意力;Zeng et al. 2022 的著名论文《Are Transformers Effective for Time Series Forecasting?》里一层线性层 DLinear 干掉了包括 Informer 在内的一票 Transformer。长序列预测的主要矛盾常常是趋势/周期的显式建模,不是注意力容量。
陷阱三:「加速是免费的」。 纯 Python 循环下 ProbSparse 的逐样本推理反而比批量矩阵乘的 full attention 慢——O(L·lnL) 的常数项(采样、排序、索引)在小 L 时吃掉理论收益。实测 L=128 时加速只有 1.1 倍。工程上小于 512 的序列别折腾稀疏注意力。
陷阱四:「M 度量精确识别活跃查询」。 M 与真实熵偏离的相关只有 0.55,且采样估计还有额外方差。排序边缘的查询会被误分类——这也是 c 不能太小的原因之一。
陷阱五:「金融序列可以直接套 Informer」。 Informer 的甜区是电力、交通这类有强周期的长序列。金融收益率的自相关近零、信噪比极低,“长序列”里多数历史对下一步没有可提取信息,注意力学到的常是噪声。金融落地更现实的路径:把注意力用在截面(股票间关系)而非超长时序上。
收尾#
ProbSparse 注意力是一个漂亮的算法工程:用 max−mean 代理稀疏度、用采样把打分成本也压到 O(L·lnL)、用均值填充懒查询——三步环环相扣,L=4096 实测 53 倍加速是实打实的。但我们的实验同样清楚地展示了它的边界:注意力不长尾时,稀疏化就是精度屠杀(c=10 时 R² 从 0.746 崩到 −0.40)。用它之前,先把你训练好的注意力矩阵可视化一下——长尾是前提,不是信仰。