尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Mixture-of-Depths Attention——深度混合注意力

Mixture-of-Depths Attention——深度混合注意力 一、研究背景与问题1.1 核心问题现代大语言模型LLM的性能提升主要依赖深度扩展增加层数但随着网络加深出现两大难题优化困难Optimization Difficulty梯度传播受阻信息稀释Information Dilution浅层形成的有用特征在逐层残差更新中被“叠加噪声”冲淡深层难以恢复1.2 现有方案不足残差连接ResNet风格结构简单但持续压缩深度历史无法解决信息稀释密集跨层连接DenseNet风格保留所有历史状态信息无损但参数和计算量随层数呈O(L²D²)增长在大模型规模下不可行因此核心目标是在深度扩展的同时有效利用历史深度信息且保持硬件效率和可控开销。二、主要创新MoDA机制2.1 核心思想将“深度历史信息”作为一种可检索的资源与标准的序列注意力融合到一个统一的Softmax操作中。每个注意力头不仅关注当前层的序列KV对标准注意力还关注同一查询位置来自所有先前层的深度KV对跨层深度记忆所有注意力分数在同一个Softmax函数下联合归一化形成统一的表示空间2.2 概念框架“读取-操作-写入”论文用此框架系统对比了三种深度流利用方式机制读取方式写入方式特点深度残差恒等映射加法叠加简单但信息稀释严重深度密集线性投影所有历史状态沿深度拼接信息无损但开销巨大深度注意力过渡用注意力自适应读取历史深度KV拼接当前层KV数据依赖成本降低MoDA最终方案融合序列深度注意力联合Softmax拼接当前层KV含FFN的KV最参数高效硬件友好2.3 关键技术细节深度KV来源不仅来自注意力层还通过轻量投影为FFN层生成对应的深度KV实验证明FFN的深度信息贡献显著统一Softmax将序列KV和深度KV拼接后共同计算注意力权重使模型能动态决定从序列还是深度历史中获取信息三、硬件高效实现核心工程贡献MoDA的朴素实现存在非连续内存访问问题GPU效率极低。论文提出了三层次优化3.1 兼容Flash的深度KV布局将深度缓存展平为长度T × L的一维数组使深度查找变为连续块读取适配FlashAttention风格内核3.2 分块感知Chunk-aware深度KV布局避免每个查询块扫描全局 T×L 深度轴而是只访问其覆盖的局部 C×L 区域深度利用率从1/T提升到1/C3.3 组感知Group-aware索引利用GQA中 G 个查询行共享同一基础时间索引的特性进一步将有效深度跨度缩减为(C/G)×L深度利用率理论上达到 100%G/G3.4 效率成果融合内核在64K序列长度下达到FlashAttention-2 效率的 97.3%相比原生PyTorch实现加速约1458倍在长序列场景下深度路径的额外开销被序列计算充分摊销额外时间从25.86%降至2.73%四、实验验证4.1 训练设置模型规模700M 和 1.5B训练数据OLMo2数据集的400B令牌采用GQA序列长度40964.2 主要实验结果模型下游任务平均性能C4验证PPLOLMo2 700M57.1118.59MoDA 700M58.87 (1.76)18.21 (-0.38)OLMo2 1.5B62.2816.16MoDA 1.5B64.39 (2.11)15.97 (-0.19)计算开销仅增加3.7% FLOPs10个下游任务全部一致提升涵盖常识推理HellaSwag、WinoGrande、科学问答ARC、SciQ、广泛知识MMLU、BoolQ10个领域验证集C4、ICE、Pile、Reddit等PPL全部降低4.3 关键消融结论深度KV本身仅重用量用前一层的序列KV作为深度KV几乎零成本即可显著提升性能FFN的深度KV额外为FFN生成深度KV带来最佳的精度-效率权衡额外注意力KV投影几乎饱和收益微小不建议使用Post-Norm优于Pre-Norm在深层模型中Post-Norm配合MoDA收益更大五、深入分析5.1 层数实验在24层和48层模型上MoDA均一致降低验证损失48层Post-Norm下MoDA比普通注意力损失降低0.05783.4062 → 3.34845.2 注意力可视化中间层和深层对深度KV块分配了显著且持续的注意力质量MoDA改变了典型的“注意力沉没”现象——注意力不集中于少数固定位置而是更广泛地分布在有用信息上5.3 效率消融优化层级运行时间ms加速比原生PyTorch2128.91× Flash兼容布局13.1162.5× 分块感知6.3338× 组感知索引1.461458×六、讨论与未来方向6.1 工业级扩展当前内核已达研究级高效但面向万亿参数工业训练仍需更深入的CUDA优化内存调度、计算流水线、通信重叠6.2 有界深度KV槽缓存当深度极大时缓存所有历史深度KV会成为内存瓶颈提出固定大小深度KV槽缓存S L通过动态选择或滑动窗口策略保留最重要的深度记忆将无界缓存变为有界缓存内存开销从深度依赖变为槽依赖七、总结MoDA的核心贡献在于算法层面将深度历史信息纳入统一的注意力框架以数据依赖方式动态检索有效缓解信息稀释工程层面通过Flash兼容布局、分块感知和组感知索引实现硬件高效的融合内核几乎不牺牲长序列训练速度实验层面在700M~1.5B规模、400B令牌训练下一致且显著优于强基线验证了深度感知聚合作为深度扩展基础机制的有效性论文最后强调MoDA是架构无关的可推广至多模态、视觉理解、世界模型等Transformer应用领域并期待其成为开源社区构建更强模型的基石。这里是自己的论文阅读记录感兴趣的话可以参考一下如果需要阅读原文的话可以看这里如下所示项目地址在这里如下所示摘要扩展深度是推动大语言模型LLM发展的关键因素。然而随着LLM变得更深它们常常遭受信号退化Signal Degradation的问题在浅层形成的信息特征会被重复的残差更新逐渐稀释使得它们在更深层中更难以恢复。我们引入了深度混合注意力Mixture-of-Depths AttentionMoDA机制该机制允许每个注意力头关注当前层的序列键值对KV对以及来自前面各层的深度键值对深度KV对。我们进一步描述了一种针对MoDA的硬件高效算法该算法解决了非连续内存访问模式的问题在序列长度为64K时达到了FlashAttention-2效率的97.3%。在1.5B参数模型上的实验表明MoDA在强基线模型上表现出一致的优越性。值得注意的是它在10个验证基准上将平均困惑度Perplexity提高了0.2在10个下游任务上将平均性能提高了2.11%而计算开销FLOPs仅为3.7%。我们还发现将MoDA与后归一化Post-Norm结合使用比与前归一化Pre-Norm结合使用效果更好。这些结果表明MoDA是深度扩展的一种有前景的基础机制。图2在1.5B参数设置下比较MoDA与强大的开源基线模型OLMo2 [27] 的验证损失和下游性能。使用MoDA的模型在C4 [30] 验证损失和下游性能即HellaSwag [48]、WinoGrande [32] 和 ARC-Challenge [10]上均优于OLMo2。1 引言近年来大语言模型LLMs[1 15, 23 37]的进展主要由四个主要维度的扩展驱动上下文长度 [8, 11 47]、训练数据 [1, 37]、模型宽度 [4, 38] 和模型深度 [6, 40]。尽管这些维度仍然有效但增量收益的成本越来越高这激发了对补充性架构扩展策略的兴趣。在当前的大语言模型实践中扩展通常更多地通过数据、上下文尤其是宽度来实现这些维度的优化行为和系统效率在规模化时通常更容易实现。相比之下深度尽管具有强大的表示潜力但仍未得到充分利用。原则上更深的堆栈可以支持更丰富的层次化计算。然而由于优化问题 [16] 和信息稀释 [20 28]现代Transformer往往无法将额外的层转化为相应的性能提升。由此产生的问题是架构设计的核心模型如何在扩展深度的同时保持优化稳定性并防止信息稀释标准的残差路径ResNet风格改善了深度网络 [16] 中的优化稳定性但它仍然将深度历史压缩到单一的隐藏状态轨迹中使得信息稀释问题在很大程度上未得到解决。许多方法 [22 42, 49] 尝试通过改进残差连接来解决这个问题。密集跨层连接DenseNet风格保留了更丰富的层级历史信息从而缓解了信息稀释 [7, 20 28]但其参数增长在大语言模型规模下是巨大的这限制了其作为主流架构的采用。注意力机制 [39] 在序列建模中的成功表明了一个更广泛的原则数据依赖的动态混合可以比固定模式的聚合更有效地保留和检索历史信息。这促使我们将相同的原则从序列建模扩展到深度建模即使每一层能够自适应地从更早的层中读取有用的状态。因此自适应跨层检索很有前景但在实际设计中仍需在表现力、效率和硬件友好性之间取得更好的平衡。在这项工作中我们引入了深度混合注意力MoDA这是一种统一的注意力机制其中每个头同时关注当前层的序列KV对和所有先前层的深度KV对。在方法论上我们通过“读取、操作、写入”Read, Operate, Write的视角分析了Transformer堆叠在共同的设计空间中比较了深度残差、深度密集和深度注意力。MoDA占据了一个高效的点它在没有密集跨层开销的情况下保留了数据依赖的深度检索能力。为了使MoDA在实际规模下具备实用性我们开发了一种硬件感知的实现 [12, 13 43]它在一次前向传播中融合了序列注意力和深度注意力并共享在线Softmax状态。此外所提出的分块感知深度KV布局Chunk-aware Depth-KV Layout和组感知索引Group-aware Indexing显著提高了内存访问效率。这个融合后的内核在64K序列长度下达到了FlashAttention-2效率的97.3%表明深度感知聚合可以在不牺牲现代GPU效率的情况下集成。我们在仅解码器Decoder-only语言模型上验证了MoDA这些模型使用OLMo2配方 [27] 在700M和1.5B规模下于400B令牌的数据上进行了训练。在我们主要的1.5B设置中MoDA在10个验证基准上将平均困惑度提高了0.2在10个下游任务上将平均下游性能提高了2.11%。我们还发现将MoDA与后归一化Post-Norm结合使用比与前归一化Pre-Norm结合使用效果更好。其他分析如模型大小扩展、注意力可视化和层数研究显示了稳健的性能提升并通过更好地将概率分配给信息丰富的序列和深度KV对减少了注意力沉没Attention Sink[41] 现象。本文的贡献总结如下我们提出了MoDA一种用于序列和深度动态混合的统一注意力公式它改进了深度信息的聚合并以数据依赖的方式解决了现代大语言模型的信息稀释问题。我们提出了一种硬件高效的融合算法使得MoDA能够应用于长上下文大语言模型训练。在64K序列长度下它达到了FlashAttention-2效率的97.3%且数值精度在允许范围内。我们提供了大量的经验证据和全面的消融实验表明MoDA在多个模型规模的大规模语料库上始终且显著地优于强大的开源基线模型OLMo2验证了每个设计选择并将MoDA确立为LLM深度扩展的可靠基础。2 深度混合注意力2.1 预备知识2.2 沿深度流堆叠Transformer深度神经网络在多个领域取得了突破性进展尤其是在引入残差连接 [16] 之后。扩展研究 [18, 19, 21] 进一步表明增加深度可以显著提高性能 [33 36]。这引出了一个自然的问题残差连接是沿深度流传播信息的最优机制吗沿着深度流我们可以将Transformer块视为一个三步过程读取、操作和写入。我们使用这个视角来描述堆叠Transformer块的不同机制。为清晰起见前两种机制深度残差 [16] 深度密集 [20, 28]是用于定义深度流设计空间的参考设计。我们引入深度注意力作为一种中间公式和概念桥梁。我们在本节的主要技术贡献从深度混合注意力MoDA开始它将序列和深度检索统一在一个Softmax算子中。图3利用深度流的机制的概念性比较。(a) 深度残差 [16] 是沿深度的标准残差连接它读取当前表示并通过加法写回。(b) 深度密集 [20 28] 读取一组历史表示并将其线性投影回宽度 DD它通过沿深度拼接来写回保留所有中间状态。(c) 我们引入深度注意力作为一种中间公式它使用注意力以数据依赖的方式读取历史深度KV对。它通过沿深度拼接当前层的键和值来写回。(d) 我们提出了深度注意力的升级版本即深度混合注意力MoDA它将深度注意力与标准序列注意力相结合。它将当前层的输出及其KV对都写入深度流以供后续层使用。3 硬件感知的高效MoDA使用原生PyTorch [29] 实现的MoDA需要对历史深度状态进行非连续读取这会降低GPU利用率。我们开发了一种硬件感知的实现通过重组深度流张量来实现连续内存访问和融合计算。图4MoDA深度缓存访问的硬件视图。左图兼容Flash的硬件高效MoDA为每个序列维护一个长度为 T×L 的深度KV缓存因此每个查询可能扫描一个较长的拼接深度KV。右图分块感知Chunk-aware的MoDA按块大小 C 分组查询并按块重组深度KV将有效的深度跨度从每个块的 T×L 减少到 (C×L)/G其中 G 是GQA组数。这种布局提高了深度KV的计算效率并减少了内存访问开销。3.1 预备知识现代GPU针对面向吞吐量的大规模数据并行工作负载进行了优化其中相同的操作被并行应用于许多元素 [12, 13, 44-46]。因此高效的注意力内核应该被组织为暴露规则的、大规模并行的计算而不是不规则的元素级控制流。流式多处理器SMs。NVIDIA GPU由许多SM组成SM是用于并行执行和资源管理的基本片上单元。高利用率需要足够多的独立块来保持许多SM处于活动状态。在大语言模型LLM训练中当序列上下文长且批量大小相对较小时沿时间维度的并行化尤为重要。计算单元CUDA核心 vs. 张量核心。在每个SM内部指令被分派到不同的执行单元。CUDA核心支持通用算术指令而张量核心为结构化的矩阵乘加运算提供更高的吞吐量。因此实用的高性能内核应最大化规则的矩阵乘法风格的计算以更好地利用张量核心。内存层次结构HBM和片上SRAM。端到端的性能由计算吞吐量和数据移动共同决定。HBM提供大容量但访问延迟较高而片上SRAM结构即寄存器、共享内存和缓存速度快得多但容量有限。因此一个关键的设计原则是改进分块和数据重用以便热数据保留在片上并最小化HBM流量。这些原则直接激励了我们的硬件感知MoDA设计。我们重组了深度KV布局并融合了计算以减少非连续内存访问并提高有效计算利用率。3.2 MoDA的硬件感知考量3.3 硬件高效的MoDA实现表2硬件高效的MoDA与FlashAttention-2 Triton内核在“前向和反向”设置下的效率比较。我们报告了在三种扩展设置下的运行时间毫秒、深度利用率ηdepth​和相对额外时间。这里B 表示批量大小d 表示头维度C 表示块大小。所有实验均在A100 GPU上使用bfloat16数据类型进行。3.3.1 效率比较4 实验在本节中我们通过在大语言模型LLM上的实验来展示所提出的MoDA的表现力和效率。表3不同深度混合注意力MoDA变体在训练集、C4验证集和下游基准上的性能。我们在400B令牌上训练700M模型。对于MoDA设置‘序列KV’表示每个令牌仅关注序列键/值可视为普通注意力机制。‘深度KV’表示每个令牌关注其深度键/值。‘额外FFN KV投影’表示进一步将FFN的输入 X 投影到深度键/值然后用于后续的注意力操作。‘额外注意力KV投影’表示设置独立的深度键/值投影而不是重用序列注意力的原始键/值投影。宽度 D、GQA组大小 G、序列长度 T 分别设置为1024、2和4096。我们进一步报告了模型的参数量和FLOPs。4.1 实验设置模型架构与训练设置。我们在不同大小的语言模型上进行了主要实验700M和1.5B。遵循通用实践我们对700M和1.5B模型采用分组查询注意力GQA[2]。我们在OLMo2 [27] 数据集的400B令牌子集上训练它们。所有模型都使用bfloat16bf16精度进行训练。全局批量大小设置为1024上下文序列长度设置为4096。更详细的训练配置如学习率调度、AdamW [25] 优化器等遵循OLMo2 [27] 的实现。评估细节。我们在流行的基准上评估模型包括PiQA [5]、HellaSwag [48]、WinoGrande [32]、OpenBookQA [26]、BoolQA [9]、SciQA [3]、COPA [31]、MMLU [17]、ARC-easyARC-E和ARC-challengeARC-C[10]。我们进一步报告训练困惑度PPL、C4验证困惑度Val PPL以及在C4 [30]、ICE [27]、m2d2-s2orc [24]、Pile [14]、Wiki-text [27] 和dolma [34] 验证集上的各领域验证困惑度后者包括Books、Common Crawl、peS2o、Reddit和Stack。4.2 主要结果4.2.1 MoDA变体我们首先比较了不同深度混合注意力MoDA变体在700M模型大小上的结果。所有模型都使用一个调度器在2k训练步骤中预热到最大学习率3e-4然后按照余弦调度衰减到3e-5。我们在表3中展示了实验结果。为了提供公平的比较我们补充了普通注意力机制OLMo2作为基线第1行。由于额外的FFN KV投影引入了额外的参数我们还报告了具有两个额外层的更多参数基线第2行。这些方法引入了与所提出的MoDA模型相当的参数量/FLOPs。从表3中我们可以观察到i深度KV显著提高了性能。我们的方法第3行与基线第1行保持相同的参数量但将每个令牌的深度KV插入到注意力计算中。注意我们直接重用前一层的序列KV作为深度KV这不会引入额外的投影参数。仅需0.12%的额外FLOPs它就提高了0.41的训练PPL、0.11的C4验证PPL和1.17的下游平均指标第1行 vs. 第3行。iiFFN层的深度KV很重要。第3行的实验仅考虑将前一注意力层的KV视为深度KV忽略了FFN层。我们进一步添加额外的KV投影来增强原始FFN将FFN的输入 X 投影到其对应的深度键/值。比较第3行和第4行我们可以观察到结合来自FFN的KV改善了0.18的训练PPL、0.27的C4验证PPL和0.77的下游平均指标。而将第4行与更多参数基线第2行比较它提高了0.37的训练PPL、0.10的C4验证PPL和1.76的下游平均指标。值得注意的是第4行与第2行具有相似的参数量/FLOPs但取得了更好的性能这表明FFN的深度信息也对深度混合注意力MoDA有所贡献。iii额外的注意力KV投影过于饱和。基于第4行我们进一步引入了额外的深度KV投影专门将注意力层的输入 X 投影到深度键/值。比较第4行和第5行我们可以观察到引入额外的注意力KV投影仅改善了0.07的训练PPL、0.04的C4验证PPL和0.10的下游平均指标。然而这种修改引入了不小的开销参数从705.7M增加到742.4MFLOPs从8.33T增加到8.63T表明额外的注意力侧深度投影接近饱和。表5所提出的MoDA模型在不同模型大小下的各领域验证困惑度。我们在OLMo2数据集的400B令牌上训练700M和1.5B模型。宽度 D、GQA组大小 G、序列长度 T 分别设置为1024、2和4096。较低的困惑度表示更好的性能并以粗体标记。总的来说这些实验揭示了一个清晰的MoDA设计原则注入深度信息是有效的但收益对额外投影引入的位置高度敏感。特别是重用注意力侧的深度KV已经能在几乎零成本的情况下提供强大的改进而添加FFN侧的深度KV则能提供最佳的精度-效率权衡。相比之下引入额外的注意力KV投影只能带来微小的收益却伴随着显著的参数/FLOPs开销。因此我们在接下来的规模扩展实验第4.2.2节中采用第4行的设置作为默认的MoDA变体。4.2.2 扩展MoDA的模型大小我们研究了在相同的400B令牌训练预算下将模型大小从700M扩展到1.5B时MoDA的收益是否持续。我们在表4中报告了下游基准结果在表5中报告了领域级别的验证困惑度。从这两个表中我们可以观察到iMoDA在不同模型规模的下游基准上提供了稳定的平均增益。对于700M模型表4中的第1行与第2行相比平均值从57.11提高到58.87即1.76。对于1.5B模型第3行与第4行相比平均值从62.28提高到64.39即2.11。ii下游收益广泛分布在常识推理、因果推理和广泛知识任务上。在常识和因果辨别任务如HellaSwag、WinoGrande和COPA上700M模型的增益分别为 0.42、4.89 和 5.001.5B模型的增益分别为 0.38、2.37 和 4.00。在面向科学和更难推理的任务如OpenBookQA、ARC-C和SciQ上700M模型的增益分别为 1.60、1.34 和 0.101.5B模型的增益分别为 2.80、4.35 和 1.50。我们也在广泛知识基准包括BoolQ增益分别为 3.09 和 3.73以及MMLU增益分别为 0.92 和 1.86上观察到了一致的增益。iii验证困惑度的收益在各领域广泛且一致。在表5中700M模型的第1行与第2行相比平均PPL从15.61降至15.46所有十个领域都有所改善。最大的700M模型降幅出现在m2d2-s2orc上PPL从24.37降至23.64。在1.5B规模下第3行与第4行相比平均PPL从13.67降至13.47同样改善了所有十个领域。显著的1.5B模型降幅出现在Reddit从21.21降至20.85、ICE从15.37降至15.08和Wiki-text从10.41降至10.16上。总的来说这两个表从互补的评估视角提供了一致的证据。表4显示了端任务性能的提升而表5显示了跨不同领域的语言建模质量的提升。表6在更深48层和更浅24层模型设置下MoDA的层数分析。我们比较了普通注意力OLMo2和具有不同MoDA选择的MoDA变体在Pre-Norm和Post-Norm配置下。模型使用相同的数据配方进行训练我们报告参数量、FLOPs和FineWeb-Edu验证损失。在这两种深度设置下引入深度KV都能持续改善验证损失而添加额外FFN KV投影能在适中的计算开销下带来进一步的收益。4.3 分析4.3.1 层数对MoDA的影响分析为了研究MoDA在不同深度预算下是否仍然有效我们使用FineWeb-Edu数据流程在小型模型上进行了层数实验。我们从FineWeb-Edu中保留了一个额外的留出分割用于验证并报告所有设置的验证损失。具体来说我们评估了更深模型48层和更浅模型24层并在Pre-Norm/Post-Norm配置下比较普通注意力与MoDA变体。在本小节的所有运行中模型宽度为384查询头数量为6键/值头数量为2。从层数实验结果中我们观察到i深度KV在不同层数下均能持续改善验证损失。对于48层模型在Pre-Norm设置下添加深度KV将损失从3.3800降至3.3759第1行 vs. 第3行在Post-Norm设置下损失从3.4062降至3.3653第2行 vs. 第4行。对于24层模型添加深度KV也将损失从3.4740降至3.4537第7行 vs. 第8行。ii在更深模型中Post-Norm从深度KV中获得的收益比Pre-Norm更大。在48层时第2行与第4行相比Post-Norm的损失降低了0.0409而第1行与第3行相比Pre-Norm的损失仅降低了0.0041。这表明对于更深的堆栈深度KV在Post-Norm配置中具有更强的优化影响。iii在深度KV的基础上额外FFN KV投影能提供进一步的收益。对于48层模型在Pre-Norm下添加额外FFN KV投影将损失从3.3759进一步降至3.3656第3行 vs. 第5行在Post-Norm下损失从3.3653降至3.3484第4行 vs. 第6行。对于24层模型它进一步将损失从3.4537降至3.4338第8行 vs. 第9行。总的来说这些结果表明MoDA在层扩展下仍然有效并且在计算预算允许时FFN侧的深度信息能带来额外的收益。4.3.2 通过注意力可视化分析MoDA为了更好地理解MoDA如何改变令牌交互我们可视化了在400B令牌上训练的700M模型的注意力热图图5。在组合Softmax公式下每个查询关注拼接的序列KV/深度KV空间红色虚线表示边界。值得注意的是深度KV部分包含注意力KV和FFN KV。从热图中我们观察到在深度KV块上存在显著且持续的注意力质量尤其是在中间层和深层。这表明模型主动检索跨层深度信息而不仅仅依赖序列局部上下文。我们还发现了一种互补模式具有更尖锐对角线序列注意力的头仍然会将部分概率分配给深度槽而分布更宽的头往往更依赖深度KV条目。另一个重要观察是MoDA展现的注意力模式与在可视化头中观察到的典型注意力沉没行为不同。MoDA并没有将大部分概率质量坍缩到少数固定的沉没位置而是将注意力更广泛地分布在序列和深度槽位上包括那些可能对任务相关的槽位。这种定性差异表明MoDA可能会改变长上下文设置中注意力质量的分配方式。特别是可视化表明部分概率质量从固定的沉没位置重新分配到了可能携带有用信息的序列/深度位置。尽管这些模式很有趣但它们确切的功能作用仍不清楚需要进一步研究。总的来说可视化结果与MoDA的核心直觉一致深度信息可以作为标准序列注意力的补充检索通道。同时改变的注意力沉没模式可能指向原始设计动机之外的其他机制或见解这应在未来的工作中更仔细地研究。4.3.3 通过效率分析MoDA为了量化每个内核设计对实际效率的贡献我们进行了增量消融实验并在表7中报告了端到端的“前向和反向”运行时间。所有实验均在单个A100 GPU上使用bfloat16进行固定设置为 B1T1024G8Hq​64Hk​8d64L64C64。注意原生PyTorch实现未针对效率进行优化我们仅在短序列长度即 T1024下报告比较结果。从表7中我们观察到i兼容Flash的深度KV布局已经比原生实现提供了数量级的加速。第1行与第2行相比运行时间从2128.900毫秒降至13.102毫秒即大约快了 162.5×。ii分块感知的深度KV布局通过减少内存访问开销进一步提高了效率。图5采用组合Softmax公式的深度混合注意力MoDA热图。列对应于均匀采样的层 {0,11,23,35}行对应于每层中随机选择的头。第一列显示仅对序列KV的注意力而其他列显示拼接的序列KV/深度KV红色虚线标记了两个KV块之间的边界。在不同层和头中深度KV块被持续分配了大量注意力表明MoDA除了标准的序列注意力外还有效地利用了深度信息。在兼容Flash的基础上第2行与第3行相比运行时间从13.102毫秒降至6.286毫秒减少了 52.0%。iii组感知索引对于充分利用组重用机制至关重要。添加组感知索引第3行 vs. 第4行进一步将运行时间从6.286毫秒降至1.460毫秒提供了额外的 4.31× 加速。总的来说结合所有三种优化方式可以获得最佳运行时间并且相比于原生PyTorch基线第1行 vs. 第4行实现了约 1458× 的端到端加速。5 结论在本文中我们提出了MoDA一种用于大语言模型的统一深度感知注意力机制旨在改进深度信息聚合并缓解由优化困难和信息稀释导致的深度效率差距。我们进一步开发了一种硬件感知的融合内核该内核具有统一的在线Softmax状态、分块感知的深度KV布局和组感知索引以维持高效的长上下文执行。在700M和1.5B模型上使用OLMo2配方进行的实验表明在适中的开销下模型在困惑度和下游性能方面均取得了一致的提升。这些结果表明显式检索历史深度信息是扩展Transformer深度的一种实用且有效的基础机制。我们将发布MoDA的完整实现并希望它能作为开源社区构建更强大语言模型的基础。除了语言建模MoDA与架构无关可以轻松集成到多模态智能、视觉理解和世界模型中这些领域正越来越多地采用Transformer。我们相信原则性的深度感知信息聚合将为这些不同领域带来广泛而持久的益处。6 讨论6.1 通过高级CUDA工程扩展MoDA以适应工业训练尽管当前的硬件感知MoDA内核已经实现了与FlashAttention 2竞争性的效率但它并非工业规模训练例如万亿参数模型的终点。在大型生产运行中额外的CUDA工程仍然至关重要包括改进的内存调度、更深层的计算流水线化以及融合注意力内核与分布式通信之间更紧密的重叠。这些优化不会改变MoDA的算法行为但可以进一步减少内存停顿和内核启动开销提高端到端吞吐量并提升集群级训练效率。因此我们将未来的CUDA优化视为一个重要的方向以将MoDA从一个高效的研究算子转变为工业LLM训练的稳健基元。6.2 通过有界深度KV槽缓存缓解内存瓶颈当扩展到非常深的网络时缓存来自所有历史层的所有深度KV状态会引入大量的内存和带宽开销。该成本随深度线性增长并可能成为长上下文训练和服务中的主要瓶颈。因此全量深度KV缓存在工业规模下越来越难以维持。一个实用的方向是使用固定大小的深度KV槽缓存。与存储所有深度KV条目不同每个查询只关注一个有界的槽集。槽预算固定为 S其中 S≪L系统动态决定保留哪些深度KV条目。两种自然的策略是动态选择和滑动窗口。动态选择根据效用评分候选深度KV条目并保留前 SS 个条目。滑动窗口策略则保留最近的深度KV条目并淘汰较旧的。也可以采用混合设计其中部分槽保留给最近性其余部分保留给高分全局记忆。这种设计将有效的深度记忆从无界缓存变为有界缓存。内存和带宽项从深度依赖扩展变为槽依赖扩展。它也为融合内核实现提供了稳定的张量形状。在实践中关键的挑战是槽分配的质量。未来的工作应研究如何与MoDA联合训练选择策略以及如何在固定的槽预算下平衡质量、延迟和硬件效率。
返回列表