跳到主要内容
返回时间线
arXiv来源发表:

稀疏注意力理论分析:JS散度由截断质量唯一决定,且随序列变长信息损失增大、泛化界更紧

相关研究与后续进展

核心概要

该工作对稀疏注意力做统一理论分析:证明全注意力与稀疏注意力之间的Jensen–Shannon散度仅由被丢弃的截断质量决定并给出闭式表达,在注意力分数为独立同分布次高斯假设下用次序统计量证明截断质量的高概率下界随序列长度增长而增大,同时用Rademacher复杂度与Xu–Raginsky互信息加熵和覆盖数分析给出随稀疏度收紧的泛化界,并在GPT-2(124M)与WikiText-103上验证闭式散度与信息损失随序列增长的趋势。

Source-provided article image: On the Trade-off Between Information Loss and Generalization in Sparse Attention

(a) M = 𝒪 ⁡ ( L ) M=\mathcal{O}(\sqrt{L})

arXiv

深度剖析

论文给出全注意力分布与稀疏注意力分布之间JS散度的闭式表达,证明该散度仅是截断质量的函数,且该函数在[0,1]上严格凸、严格递增。 此前对稀疏注意力的理论讨论多停留在图论属性或经验观察,缺少把近似误差归结为单一标量的定量刻画;这里把信息损失完全压缩到一个可解释的截断质量上。 基于权重缩放引理与KL散度逐项计算的解析推导,并在GPT-2(124M)上对WikiText-103测试集注意力行做top-截断验证,经验值与闭式值在float32精度内一致。

在注意力分数为独立同分布次高斯、投影谱范数与嵌入范数有界的假设下,论文用次序统计量给出保留与丢弃质量之比的高概率上界,从而得到截断质量的高概率下界随序列长度增大而增大的结论。 把信息损失随序列长度的增长从经验现象提升为有概率保证的定量陈述,并指出其速率由稀疏预算阶数与由嵌入尺度、投影谱范数、头维度共同决定的系数控制,而非仅由稀疏度决定。 解析证明结合次高斯尾界与次序统计量集中分析;实验上六种稀疏策略的截断质量均随序列长度上升,方向与定理一致,但界比实测保守一至三个数量级。

论文用Dudley积分与体积论证给出M-稀疏注意力层的Rademacher复杂度上界,并据此得到固定稀疏模式下的泛化界,指出稀疏度降低假设类复杂度、使泛化界更紧。 这是对稀疏注意力Rademacher复杂度的形式化刻画,并把单层结果通过Lipschitz常数推广到多层,把泛化间隙与逐层Lipschitz常数和假设容量的乘积联系起来。 解析推导;实验用高斯噪声随机化测试作为Rademacher复杂度的经验代理,观察到稀疏度增大时泛化间隙变宽,与理论增长曲线方向一致。

论文结合Xu–Raginsky互信息界与稀疏假设类的熵分解和覆盖数分析,得到同时依赖稀疏度与权重离散分辨率的泛化界,把结构复杂度与数值精度纳入同一框架。 将信息论泛化分析与稀疏注意力的结构性质连接起来,用组合项刻画支撑集结构复杂度、用几何项刻画稀疏权重分辨率,与JS散度形成互补视角。 解析推导,明确限定在固定、与输入无关的支撑集与限制在单纯形ε-网上的稀疏权重假设下成立。

启示与展望

该分析适用于固定、与输入无关的稀疏支撑集设定,作者将其定位为隔离稀疏本身效应的基线复杂度界,可与自适应路由引入的额外复杂度解耦。对采用长序列的Transformer部署者,结论给出可操作指引:稀疏预算阶数与由嵌入尺度、投影谱范数、头维度共同决定的系数共同控制信息损失累积速度,在固定嵌入维度下减小头维度(即增加头数)可减缓截断质量随序列长度的增长,但受注意力熵塌缩约束限制。JS散度闭式可用于在给定截断质量下直接估计近似误差,而熵与覆盖数界则把权重离散分辨率纳入泛化评估,适合量化实现场景。

理论界建立在注意力分数独立同分布次高斯、投影谱范数与嵌入范数有界等假设上,作者也指出i.i.d.次高斯模型未刻画真实注意力在少数键上的强集中,因此界保守而非紧,与实测相差一至三个数量级。数据依赖稀疏模式(如Top-k或路由式)下的假设类随输入变化,标准Rademacher分析不能直接套用,作者将其列为未来方向。泛化实验以高斯噪声随机化测试作为Rademacher复杂度的经验代理,其与真实数据分布下的泛化间隙关系仍需进一步观察。此外,加载文本中多处公式与图注以占位形式呈现,具体常数与曲线细节无法从文本直接核对,读者若需精确数值应查阅原文图表。

来源