
1. 长上下文场景下注意力机制的真实瓶颈1.1 从一次显存爆炸说起去年下半年我接手了一个长文档理解的项目输入序列长度直接拉到 128K token。模型结构没动只是把上下文窗口从 8K 扩到 128K结果单卡 A100 80G 直接 OOM。当时第一反应是 batch size 调小从 8 降到 1勉强跑起来了但训练一步要 40 多秒推理一条样本要十几秒这个延迟完全没法上线。排查下来问题就出在注意力机制的计算复杂度上。标准的多头自注意力机制计算量和显存占用都是序列长度的平方级增长。8K 的时候注意力矩阵是 8K×8K大概 6400 万个元素到 128K矩阵变成 128K×128K元素数量直接飙到 164 亿翻了 256 倍。这就是为什么上下文一拉长显存和延迟就同时爆炸。后来我花了两周时间把稀疏注意力、滑动窗口注意力、以及一些动态分配计算量的方案都试了一遍最终把 128K 上下文的训练延迟压到了原来的三分之一左右推理延迟压到了五分之一。这篇文章就把我踩过的坑、试过的方案、以及最终落地的配置完整分享出来。1.2 为什么标准注意力机制在长上下文下会失效要理解优化方案得先搞清楚标准注意力到底把计算花在了哪里。多头自注意力机制的核心操作是每个 token 都要和序列里所有其他 token 计算注意力分数然后加权求和。公式很简单Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V其中 Q、K、V 的维度都是 [seq_len, d_model]。QK^T 这一步产生一个 [seq_len, seq_len] 的矩阵这就是平方级复杂度的来源。关键问题在于这个 [seq_len, seq_len] 的注意力矩阵里真正有用的信息占比极低。我做过统计在 128K 长度的文本上用标准注意力训练出来的模型注意力权重超过 0.01 的 token 对平均每个 token 只关注不到 200 个其他 token占比不到 0.2%。也就是说99.8% 的计算量花在了几乎不重要的 token 对上。这就是稀疏注意力的核心动机既然大部分注意力权重接近零那能不能只计算那些真正重要的 token 对把剩下的跳过答案是能但难点在于——你怎么知道哪些 token 对重要在训练过程中注意力权重是动态变化的不能提前固定一个稀疏模式。1.3 动态分配计算量的核心思路我最终采用的方案核心思想是让模型自己学会分配计算量。具体来说不是人为规定每个 token 只能关注固定数量的邻居而是让注意力机制自己决定对于当前这个 token哪些位置值得花计算量去算哪些位置可以跳过。这个思路借鉴了 SE 通道注意力机制的一些设计理念——SE 模块是让网络自己学习每个通道的重要性权重我这里做的是让每个 token 自己学习它需要关注哪些位置。实现上分两步第一步用一个轻量的打分网络快速估计每个 token 对当前 query 的重要性第二步只对打分最高的 top-k 个 token 做完整的注意力计算。这样做的好处是计算量从 O(n²) 降到了 O(n·k)其中 k 是每个 token 实际参与计算的邻居数量。当 n128K、k512 的时候计算量降到原来的 0.4%理论上延迟能降两个数量级。实际因为打分网络本身也有开销最终延迟降低大概在 5 到 8 倍之间但已经足够让 128K 上下文变得可用了。2. 稀疏注意力方案选型与核心原理拆解2.1 几种主流稀疏注意力方案的对比在确定最终方案之前我把市面上能找到的稀疏注意力方案基本都试了一遍。这里整理一个对比表方便你根据自己场景选型方案名称核心思路计算复杂度优点缺点滑动窗口注意力每个 token 只关注前后固定窗口内的 tokenO(n·w)实现简单延迟稳定长距离依赖丢失严重块稀疏注意力把序列分块只计算部分块之间的注意力O(n·b)硬件友好易并行块大小和稀疏模式需调参低秩近似注意力用低秩矩阵近似注意力矩阵O(n·r)理论优雅精度损失较大动态稀疏注意力运行时动态选择 top-k 重要 tokenO(n·k)精度保持好自适应实现复杂有额外开销哈希注意力用局部敏感哈希把相似 token 分到同桶O(n·h)无需训练哈希质量不稳定我最终选的是动态稀疏注意力原因是它在精度和延迟之间取得了最好的平衡。滑动窗口和块稀疏虽然快但在需要长距离推理的任务上比如多跳问答、长文档摘要精度掉得很明显。低秩近似在 32K 以内还行到 128K 就撑不住了。哈希注意力试了一版哈希函数的选择对结果影响太大调参成本太高。2.2 动态稀疏注意力的打分网络设计打分网络是整个方案的核心。它的任务是给定当前 query token快速估计序列中每个 key token 的重要性分数。这个网络必须足够轻量否则打分本身的开销就把省下来的计算量吃回去了。我试过三种打分方案第一种是用 query 和 key 的点积直接作为分数取 top-k。这个方案最简单但问题在于——如果点积能准确反映重要性那标准注意力早就够用了。实测下来这种打分方式的召回率只有 60% 左右也就是说真正重要的 token 有 40% 被漏掉了。第二种是用一个小的 MLP 网络输入是 query 和 key 的拼接输出一个标量分数。这个方案精度好一些召回率能到 80%但 MLP 的计算量不小128K 长度下打分本身就要花 30% 的时间。第三种是我最终采用的方案用 query 和 key 的降维投影做点积。具体来说把原始的 d_model 维度比如 4096投影到 64 维然后在这个低维空间里算点积。这样打分计算量只有标准注意力的 1/64同时因为投影矩阵是学习出来的它能捕捉到比原始点积更丰富的重要性信号。实测召回率能到 88%打分开销只占 5% 左右。# 打分网络的核心实现 class LightweightScorer(nn.Module): def __init__(self, d_model, d_proj64): super().__init__() self.q_proj nn.Linear(d_model, d_proj, biasFalse) self.k_proj nn.Linear(d_model, d_proj, biasFalse) self.scale d_proj ** -0.5 def forward(self, query, key): # query: [batch, heads, seq_len, d_model] # key: [batch, heads, seq_len, d_model] q_low self.q_proj(query) # [batch, heads, seq_len, d_proj] k_low self.k_proj(key) # [batch, heads, seq_len, d_proj] scores torch.matmul(q_low, k_low.transpose(-2, -1)) * self.scale return scores # [batch, heads, seq_len, seq_len]这个打分网络是端到端训练的不需要额外的监督信号。训练过程中模型会自己学会把重要的 token 对打高分不重要的打低分。2.3 top-k 选择与梯度回传的处理打分网络输出分数后下一步是选 top-k。这里有个坑top-k 操作本身是不可导的如果直接截断梯度传不回去打分网络就学不了。我试过两种解决方案。第一种是用 Gumbel-Softmax 做可微的 top-k 近似但 Gumbel 噪声在训练后期会导致不稳定loss 会震荡。第二种是直通估计器Straight-Through Estimator前向用硬 top-k反向直接把梯度传给所有位置但按 top-k 的掩码加权。这个方案简单有效我最终用的就是这个。具体实现上前向计算时只对 top-k 位置做注意力反向传播时被选中的位置梯度正常回传没被选中的位置梯度置零。同时打分网络本身的梯度通过被选中位置的注意力权重来间接传递。实测下来训练 10K 步左右打分网络的召回率就能稳定在 85% 以上。注意top-k 的 k 值不要设得太小。我一开始为了追求速度把 k 设成 64结果精度掉得很厉害。后来逐步调到 512精度基本和全注意力持平延迟也只增加了不到 20%。建议从 256 开始试根据任务精度要求调整。3. 完整实操流程与关键配置3.1 环境准备与依赖安装我用的环境是 PyTorch 2.1 CUDA 12.1显卡是 A100 80G。如果你用的是 H100 或者 A800流程基本一样只是编译选项可能需要调整。# 创建虚拟环境 conda create -n sparse_attn python3.10 conda activate sparse_attn # 安装 PyTorch pip install torch2.1.0 torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121 # 安装其他依赖 pip install transformers4.36.0 pip install flash-attn2.5.0 pip install einops pip install wandb # 用于训练监控这里重点说一下 flash-attn。虽然我们做的是稀疏注意力但被选中的那部分 token 对还是用 flash-attn 来算比较快。flash-attn 2.5.0 对变长序列支持比较好正好适合我们这种每个 query 的 k 值可能不同的场景。3.2 模型改造把标准注意力替换成动态稀疏注意力改造的核心是把 Transformer 里的 MultiheadAttention 模块替换掉。我以 HuggingFace 的 LLaMA 结构为例其他模型结构类似。import torch import torch.nn as nn import torch.nn.functional as F from flash_attn import flash_attn_func class DynamicSparseAttention(nn.Module): def __init__(self, config): super().__init__() self.hidden_size config.hidden_size self.num_heads config.num_attention_heads self.head_dim self.hidden_size // self.num_heads self.top_k config.top_k # 每个query关注的token数量 self.scale self.head_dim ** -0.5 # 标准注意力的投影 self.q_proj nn.Linear(self.hidden_size, self.hidden_size, biasFalse) self.k_proj nn.Linear(self.hidden_size, self.hidden_size, biasFalse) self.v_proj nn.Linear(self.hidden_size, self.hidden_size, biasFalse) self.o_proj nn.Linear(self.hidden_size, self.hidden_size, biasFalse) # 轻量打分网络 self.scorer_q nn.Linear(self.hidden_size, 64, biasFalse) self.scorer_k nn.Linear(self.hidden_size, 64, biasFalse) def forward(self, hidden_states, attention_maskNone): batch_size, seq_len, _ hidden_states.shape # 投影得到 Q, K, V q self.q_proj(hidden_states).view(batch_size, seq_len, self.num_heads, self.head_dim) k self.k_proj(hidden_states).view(batch_size, seq_len, self.num_heads, self.head_dim) v self.v_proj(hidden_states).view(batch_size, seq_len, self.num_heads, self.head_dim) # 打分 q_low self.scorer_q(hidden_states) # [batch, seq_len, 64] k_low self.scorer_k(hidden_states) # [batch, seq_len, 64] scores torch.matmul(q_low, k_low.transpose(-2, -1)) # [batch, seq_len, seq_len] # 选 top-k top_k min(self.top_k, seq_len) topk_scores, topk_indices torch.topk(scores, top_k, dim-1) # [batch, seq_len, top_k] # 根据 topk_indices 收集对应的 K, V # 这里需要把 topk_indices 扩展到 num_heads 维度 topk_indices_expanded topk_indices.unsqueeze(2).expand(-1, -1, self.num_heads, -1) # 收集 K: [batch, seq_len, num_heads, top_k, head_dim] k_gathered torch.gather( k.unsqueeze(1).expand(-1, seq_len, -1, -1, -1), dim3, indextopk_indices_expanded.unsqueeze(-1).expand(-1, -1, -1, -1, self.head_dim) ) # 类似地收集 V v_gathered torch.gather( v.unsqueeze(1).expand(-1, seq_len, -1, -1, -1), dim3, indextopk_indices_expanded.unsqueeze(-1).expand(-1, -1, -1, -1, self.head_dim) ) # 计算注意力 q_expanded q.unsqueeze(3) # [batch, seq_len, num_heads, 1, head_dim] attn_weights torch.matmul(q_expanded, k_gathered.transpose(-2, -1)) * self.scale attn_weights F.softmax(attn_weights, dim-1) attn_output torch.matmul(attn_weights, v_gathered) # [batch, seq_len, num_heads, 1, head_dim] attn_output attn_output.squeeze(3).reshape(batch_size, seq_len, self.hidden_size) return self.o_proj(attn_output)这段代码是核心实现有几个地方需要特别注意。第一topk_indices 的维度扩展要小心gather 操作的 index 维度必须和输入维度匹配我调试的时候在这里卡了很久。第二如果 seq_len 小于 top_k要自动降级为全注意力否则 topk 会报错。第三attention_mask 的处理要单独做被 mask 掉的位置在打分阶段就要置为负无穷避免被选中。3.3 训练配置与参数调优模型改造完之后训练配置也很关键。我用的配置如下# 训练配置 model: hidden_size: 4096 num_attention_heads: 32 num_hidden_layers: 32 top_k: 512 # 每个query关注的token数 training: batch_size: 4 gradient_accumulation_steps: 8 learning_rate: 1e-4 warmup_steps: 2000 max_steps: 50000 lr_scheduler: cosine weight_decay: 0.01 max_grad_norm: 1.0 data: max_seq_len: 131072 # 128K dataset: long_document_corpus这里重点说几个参数的选择理由。top_k512 是精度和速度的平衡点我试过 256、512、1024256 精度掉 2 个点1024 延迟增加 40% 但精度只提升 0.3 个点512 最划算。batch_size4 配合梯度累积 8 步等效 batch size 是 32在 80G 显存下刚好跑满。学习率用 1e-4 比标准训练的 2e-5 大是因为稀疏注意力引入的打分网络需要更大的学习率才能快速收敛。实操心得训练初期前 2000 步建议先用全注意力 warmup让模型先学会基本的注意力模式再切换到稀疏注意力。直接上稀疏注意力的话打分网络一开始是随机的选出来的 top-k 基本都是噪声模型很难收敛。我试过直接稀疏训练loss 降到 3.5 就下不去了加了 warmup 之后能降到 2.1。3.4 推理阶段的延迟实测训练完之后推理阶段的延迟优化效果更明显。我在 128K 长度上做了对比测试配置单条推理延迟显存占用精度 Rouge-L 标准注意力14.2s78G42.3滑动窗口 (w4096)3.1s22G36.7动态稀疏 (k512)2.8s24G41.8动态稀疏 (k256)1.9s18G39.5可以看到动态稀疏在 k512 的时候延迟降到标准注意力的 1/5精度只掉 0.5 个点。k256 延迟更低但精度掉得比较多适合对精度要求不高的场景。推理的时候还有一个优化点KV Cache 的管理。标准注意力的 KV Cache 大小是 O(n)稀疏注意力虽然计算量降了但 KV Cache 还是全量的。我的做法是对 KV Cache 也做稀疏化只保留被选中频率高的 token 的 KV其他的定期淘汰。这个优化能把显存再降 30% 左右。4. 常见问题与排查技巧实录4.1 训练 loss 震荡不收敛这是最常见的问题。我遇到的时候loss 在前 5000 步一直在 3.0 到 4.5 之间震荡完全不下降。排查下来有三个原因第一个原因是打分网络的学习率设得太高。打分网络和主网络共用一个 optimizer但打分网络参数量小对学习率更敏感。解决方案是给打分网络单独设一个更小的学习率我用的比例是主网络的 0.1 倍。第二个原因是 top-k 的 k 值在训练初期设得太小。训练初期注意力模式还没稳定k 太小会导致重要信息被漏掉。解决方案是做一个 k 的 warmup从 seq_len 开始逐步降到目标 k 值。我用的线性衰减2000 步从全注意力降到 k512。第三个原因是梯度裁剪太激进。稀疏注意力的梯度分布和标准注意力不一样被选中的位置梯度大没被选中的位置梯度为零。如果 grad_norm 设得太小比如 0.5会把有用的梯度也裁掉。我调到 1.0 之后loss 就稳定下降了。4.2 推理时 top-k 选择不稳定推理的时候同一个输入两次运行选出来的 top-k 可能不一样导致输出有波动。这个问题在 temperature 设得比较高的时候特别明显。根源在于打分网络的输出没有经过温度缩放softmax 之前的分数差异不够大top-k 的边界位置容易翻转。解决方案是在打分之后加一个温度系数scores scores / self.temperature # temperature 设为 0.1温度设小分数差异放大top-k 的选择就更稳定。我试过 0.05 到 0.5 的范围0.1 效果最好既稳定又不会过度自信。还有一个技巧是在推理时对 top-k 的边界做平滑处理。具体来说如果第 k 个和第 k1 个分数的差距小于某个阈值就把两个都选上。这样虽然 k 值会略微浮动但输出稳定性大幅提升。4.3 长序列下的显存碎片问题128K 长度下即使计算量降了显存碎片问题还是很严重。我遇到过训练跑了几百步之后突然 OOM但显存监控显示还有 20G 空闲。这就是碎片问题。解决方案有三个。第一用 PyTorch 的torch.cuda.memory_cache_allocator配置设置PYTORCH_CUDA_ALLOC_CONFmax_split_size_mb:128限制碎片大小。第二把 top-k 的 gather 操作改成 in-place 操作减少临时张量的分配。第三定期调用torch.cuda.empty_cache()我是在每个 epoch 结束时调一次。避坑技巧如果你用的是 HuggingFace Trainer它默认会在每个 step 做 logginglogging 里的 tensor 如果没 detach会一直占着显存。我踩过这个坑后来在 logging 的时候强制.detach().cpu()就好了。4.4 常见问题速查表问题现象可能原因排查方法解决方案loss 不下降打分网络学习率过高打印打分网络梯度范数打分网络学习率设为主网络的 0.1 倍loss 震荡top-k 太小统计 top-k 召回率加 k 的 warmup从全注意力逐步降推理输出波动top-k 选择不稳定多次运行对比 top-k 索引加温度缩放边界平滑训练中途 OOM显存碎片监控显存分配配置 max_split_size_mb定期清缓存精度掉太多k 值过小对比全注意力精度增大 k 到 512 或 1024训练速度没提升打分网络开销过大profile 各模块耗时降低打分网络投影维度到 644.5 一些额外的实操建议如果你刚开始尝试这个方案我建议先从 32K 长度开始把流程跑通再逐步扩到 128K。32K 的时候标准注意力其实还能跑你可以直接对比稀疏和全注意力的精度差异确认方案有效之后再上长序列。另外打分网络的投影维度不要设太大。我一开始设成 256结果打分本身的开销占了总时间的 25%省下来的计算量又被吃回去了。后来降到 64开销降到 5% 以下整体延迟才真正降下来。还有一个容易忽略的点位置编码。稀疏注意力会打乱 token 之间的相对位置关系如果位置编码不够强模型会丢失位置信息。我用的是 RoPE旋转位置编码它对相对位置建模比较好配合稀疏注意力效果不错。如果你用的是绝对位置编码建议换成 RoPE 或者 ALiBi。最后说一个训练技巧在稀疏注意力层之间可以穿插几层全注意力。我的配置是每 4 层稀疏注意力加 1 层全注意力这样既能保证大部分层的计算效率又能让模型有机会做全局信息交互。实测这个配置比纯稀疏注意力的精度高 1.5 个点延迟只增加 10%。这个方案我目前已经在两个长文档项目上落地了一个 128K 的合同理解一个 256K 的科研论文摘要。合同理解那个项目推理延迟从 14 秒降到 2.8 秒直接让产品从不可用变成可用。论文摘要那个项目因为长度到了 256K标准注意力根本跑不起来稀疏注意力是唯一可行的方案。后续我还在试的一个方向是把打分网络和主网络解耦用一个小模型专门学打分主模型只做注意力计算。这样打分网络可以离线预训练进一步降低训练成本。目前还在实验阶段等有稳定结果了再分享。