Prism Transformer: 渐进式头调度用于层级注意力处理
摘要
Prism Transformer 用渐进式头调度替代了统一的多头注意力机制,该调度在层间逐步增加头的数量,从而在不增加参数或计算量的情况下实现从局部到全局的层级结构。在124M、354M和757M三个模型规模上,它在语言建模和零样本基准测试中始终优于标准Transformer。
查看缓存全文
缓存时间: 2026/06/29 05:22
# 渐进式头调度用于层次化注意力处理
来源:https://arxiv.org/html/2606.27449
###### 摘要
多头注意力机制通常在每一层将隐藏维度在所有头上均匀分配,使得模型整个深度范围内的每个头都具有相同的子空间维度 \(dh = d_{model} / h\)。在本文中,我们将这种均匀分配识别为一个根本性的结构瓶颈:由于维度空间受限,早期层的头无法忠实地捕捉复杂的高维上下文模式。为了解决这个问题,我们提出了 Prism Transformer,这是一种新颖的架构范式,用渐进式头调度取代了静态、均匀的头配置。通过跨层单调增加头的数量,Prism Transformer 自然地建立了一个从局部到全局的表征层次结构:早期层利用数量更少、但维度异常宽的头来捕捉复杂的局部组合模式,而深层则部署大量狭窄的头来将这些模式分解为专门的语言特征。至关重要的是,这种结构变化是参数中性、计算中性的,并且不引入任何训练或推理开销,保持了与标准 Transformer 相同的权重矩阵和 FLOP 预算。在三种模型规模(124M、354M 和 757M)上,Prism Transformer 持续优于均匀基线,实现了验证损失的持续降低,并在下游零样本基准测试(包括 PIQA、HellaSwag、ARC-Easy 和 WinoGrande)上取得了持续的增益。我们的研究结果表明,非均匀的子空间分配能够释放标准 Transformer 预算内的潜在容量,从而实现模型容量的更有效利用。
## 1 引言
多头注意力(MHA)机制 Vaswani 等人 (2017) (https://arxiv.org/html/2606.27449#bib.bib1) 是 Transformer 架构的决定性组件。通过将查询、键和值投影到 \(h\) 个独立的子空间(维度为 \(d_h = d_{model} / h\))中,MHA 使得模型能够同时关注多个不同的表征模式。自原始 Transformer 以来,这种隐藏维度的划分在所有层中保持不变:标准解码器的每一层都使用相同数量的头,因此每个头都具有完全相同的维度。
这种均匀分配嵌入了一个关于表征处理的强烈隐含假设:即网络中所有深度的层都能从相同粒度的注意力子空间中同等受益。我们认为,这个假设引入了早期层实际需求与均匀调度强制其执行任务之间根本性的架构不匹配。
#### 均匀分配的问题
早期的 Transformer 层负责将原始的 token 嵌入整合成能捕捉复杂局部组合语义结构的高层表征。忠实地编码这些复杂的局部模式需要每个子空间具备充足的表征容量。在标准的均匀配置下(例如,对于 \(d_{model}=768\),使用 \(h=12\) 个头),每个早期层的头被限制在狭窄的 \(d_h=64\) 维度内。由于这个严重受限的维度空间,单个头被迫将小的子空间稀疏地分布在位置上,使得它们在物理上无法有效解析复杂的高维局部上下文特征。
随着表征在网络中传播,其结构需求也在演变。网络中层的核心功能是全局上下文整合,将这些密集的局部特征跨越长距离序列依赖关系进行聚合。最后,深层 Transformer 从广泛的整合转向专门的特性分解,提取并分离特定的细粒度句法、语义或任务导向信号。这种下游精炼目标自然适用于并行操作的大量狭窄、专注的子空间。通过强制每个阶段(包括局部组合、全局整合和细粒度精炼)都采用相同的结构配置,传统的均匀调度忽略了网络自然的局部到全局的演进轨迹。
#### Prism Transformer。
我们提出通过用渐进式头调度取代僵化的均匀配置来解决这个结构瓶颈:一个非递减序列 \((h_1, h_2, \dots, h_L)\),其中 \(h_l\) 表示第 \(l\) 层的头数量。早期层使用较少的头,给予每个独立的头一个异常宽的子空间(当 \(h_l\) 很小时,\(d_h^{(l)} = d_{model} / h_l\) 很大)。然后,头数量随网络深度单调增加,在较深层收敛到标准基线配置。
这种调度的形状在 \(d_h\) 上随着深度呈下降轨迹,形成一个逐渐变窄的棱镜,表征通过它从局部流向全局。这种设计产生了两个关键的 structural 属性。首先,它是参数中性的:因为 MHA 投影矩阵(\(W_Q, W_K, W_V, W_O\))保持其标准形状 \(\mathbb{R}^{d_{model} \times d_{model}}\),改变切片数量不会改变总参数量。其次,它是计算中性的:主要的注意力 FLOPs 在数学上对头数量保持不变。因此,Prism Transformer 完全免费地解锁了潜在的表征容量。
图1 (https://arxiv.org/html/2606.27449#S1.F1) 说明了均匀调度和 Prism 调度之间的对比。
参见图注 图 1:头调度可视化 各层注意力头分配对比。(左) 标准基线,采用均匀头(所有块上的平坦调度)。(右) Prism Transformer,采用递增的头调度(各块上的渐进阶梯式增长)。通过在更深层扩展头的数量,Prism Transformer 隐式地创建了一个单调递减的每头维度 (\(d_h\)) 轨迹。
#### 贡献
我们的主要贡献如下:
- • Prism Transformer:我们提出了一种用于仅解码器 Transformer 的渐进式头调度,它在零参数或计算开销的情况下建立了强大的局部到全局表征层次结构。
- • 一致的尺度不变增益:在三种模型规模(124M、354M、757M)上,Prism Transformer 在相同的训练计算量下实现了比均匀基线更低的验证损失。(表 2 (https://arxiv.org/html/2606.27449#S3.T2))。
- • 机制性注意力分析:我们进行了详细的逐层注意力距离分析,实验证明渐进式调度重构了网络的注意力分布。它鼓励早期宽头层进行紧密、高度局部的语义聚合,并将广泛、全局的整合转移到调度完成的中层网络(第 4 节 (https://arxiv.org/html/2606.27449#S4))。
- • 下游基准测试的持平与改进:在零样本语言基准测试(PIQA、HellaSwag、ARC-Easy、WinoGrande)上的评估证实,我们的渐进式调度能够保持或显著提高下游任务的准确性。(表 3 (https://arxiv.org/html/2606.27449#A1.T3)、图 2 (https://arxiv.org/html/2606.27449#S3.F2))。
- • 头调度设计原则:通过系统的结构消融实验,我们分离出了调控调度有效性的具体几何标准,例如维度过渡平滑度和块整合度。我们将这些发现形式化为一组紧凑的可迁移设计规则,这些规则在所有评估的参数规模上一致泛化(第 2 节 (https://arxiv.org/html/2606.27449#S2))。
## 2 Prism Transformer
### 2.1 渐进式头调度的数学表述
设一个 \(L\) 层 Transformer 具有恒定的模型维度 \(d_{model}\)。渐进式头调度定义为一个非递减整数序列 \(\mathcal{S} = (h_1, h_2, \dots, h_L)\),其中 \(h_l\) 表示第 \(l\) 层分配的注意力头数量,并受以下结构约束:
1. \((i)\) 可整除性:\(h_l \mid d_{model}\) 对于所有 \(l \in \{1, \dots, L\}\) 成立(头数量整除模型维度)。
2. \((ii)\) 单调性:\(h_1 \le h_2 \le \dots \le h_L\)(单调非递减)。
3. \((iii)\) 边界收敛:\(h_L = h_{\text{base}}\)(最终层匹配基线头数量)。
### 2.2 经验设计准则
虽然第 2.1 节中的约束定义了有效调度的空间,但在此空间内的优化需要调参。通过对候选形状进行系统的结构消融(详见附录 C (https://arxiv.org/html/2606.27449#A3)),我们形式化了两个决定调度有效性的关键设计原则:
1. \((i)\) 阶梯平滑度:突变的单步转换会降低性能。最优调度利用局部的多层巩固阶段(例如,在改变粒度前,将特定头维度维持至少 2 到 4 个连续层)来稳定表征。
2. \((ii)\) 基线保持阶段:为确保稳健的语义收敛,网络中至少一半的层(\(L/2\))必须致力于基线头数量 \(h_{\text{base}}\)。
### 2.3 硬件对齐与复杂性
Prism Transformer 的一个主要优势是其与现代计算集群的零开销集成。改变头数量不会触及参数空间:由于投影矩阵(\(W_Q, W_K, W_V, W_O \in \mathbb{R}^{d_{model} \times d_{model}}\))保持其标准的 uniform 形状,总参数严格不变。
此外,我们在选择有效调度时施加了硬件布局偏好。通过选择早期阶段头数量产生的维度是 2 的幂次(\(d_h^{(l)} \in \{256, 128\}\))的配置,生成的注意力切片自然地与 GPU Tensor Cores 的 tile 边界对齐。这避免了执行内存步长和未对齐的张量分割,保证了我们的表征增益在标准训练或推理过程中不会引入实际延迟惩罚。
### 2.4 参数与计算中性
#### 参数不变性
第 \(l\) 层标准多头注意力块的参数预算完全由查询、键、值和输出映射的投影算符决定:\(W_Q, W_K, W_V, W_O \in \mathbb{R}^{d_{model} \times d_{model}}\)。单层 MHA 模块的总参数量形式化为:
\[N_{\text{params}}^{(l)} = 4 d_{model}^2 + 4 d_{model} \cdot \mathbb{I}_{\text{bias}}\]
其中 \(\mathbb{I}_{\text{bias}} \in \{0, 1\}\) 表示是否存在偏置向量的指示变量。由于这种分配严格依赖于全局模型维度 \(d_{model}\),因此它从根本上独立于特定层的头数量 \(h_l\)。因此,Prism Transformer 在实现架构重平衡时,完全没有参数膨胀或结构足迹修改。
#### 计算复杂度
多头注意力层中每个 token 步的主要浮点运算(FLOPs)按三种不同操作缩放。设 \(T\) 表示序列上下文长度。第 \(l\) 层的计算分解如下:
1. \((i)\) 密集投影:获取 \(Q, K,\) 和 \(V\) 表示的线性映射需要 \(\mathcal{O}(T \cdot d_{model}^2)\) FLOPs。
2. \((ii)\) 注意力矩阵计算:计算内积注意力矩阵并应用 softmax 需要按头数量乘以各自子空间维度进行缩放:\(\mathcal{O}\left(T^2 \cdot h_l \cdot d_h^{(l)}\right) = \mathcal{O}\left(T^2 \cdot d_{model}\right)\)。
3. \((iii)\) 输出线性对齐:最终的投影矩阵 \(W_O\) 需要 \(\mathcal{O}(T \cdot d_{model}^2)\) FLOPs。
由于这些算法项均不依赖于 \(h_l\) 的选择,Prism Transformer 在理论 FLOP 分配上是严格计算中性的。
#### 硬件对齐
在 Prism Transformer 中,我们制定的渐进式头调度使得早期阶段头配置产生的维度是 2 的幂次(\(d_h^{(l)} \in \{256, 128\}\))。这种结构选择天然地保留了跨 GPU Tensor Cores 的硬件 tile 和内存对齐,完全避免了未对齐的张量分割或内存步长。因此,Prism Transformer 自然地匹配标准 uniform 模型的原始吞吐量和挂钟训练速度,以零计算成本提供其架构表征优势。(表 2 (https://arxiv.org/html/2606.27449#S3.T2))。
## 3 实验
### 3.1 实验设置
#### 硬件与核心实现
所有训练都在一个由 \(8 \times\) NVIDIA H100 (80GB SXM5) GPU 组成的分布式集群上执行,这些 GPU 通过 NVLink 互连。模型使用 PyTorch 构建,通过扩展 NanoGPT 框架 Karpathy (2022) (https://arxiv.org/html/2606.27449#bib.bib16),并加入了旋转位置嵌入(RoPE) Su 等人 (2023) (https://arxiv.org/html/2606.27449#bib.bib22) 和 SwiGLU 激活函数 Shazeer (2020) (https://arxiv.org/html/2606.27449#bib.bib23)。架构调整完全局限于注意力层的分割维度;所有其他宏观参数(例如,隐藏状态、层深度、优化器选择)在基线和实验设置中保持相同。
#### 数据集与分词
模型在 FineWeb 数据集 Penedo 等人 (2024) (https://arxiv.org/html/2606.27449#bib.bib12) 上进行预训练,使用 GPT-2 BPE 分词器进行分词。遵循 Chinchilla 缩放定律 Hoffmann 等人 (2022) (https://arxiv.org/html/2606.27449#bib.bib14),我们为每个模型训练 25 到 30 个 token 每参数,略微超过计算最优比率,以确保所有变体达到足够的收敛程度以进行稳健比较。
#### 训练配置
训练采用混合精度 bfloat16,在分布式数据并行(DDP)下进行,固定上下文长度为 1024 个 token。AdamW 优化器配置为 \(\beta_1 = 0.9\),\(\beta_2 = 0.95\),\(\epsilon = 10^{-8}\),权重衰减 0.1,梯度裁剪为 1.0。始终使用带线性预热(在前 2.5% 的训练 token 上进行)的余弦衰减学习率调度。架构规格和超参数总结在表 1 (https://arxiv.org/html/2606.27449#S3.T1) 中。
表 1:模型配置与训练超参数。\(h_{\text{base}}\) 表示基线(均匀)头数量;Prism 调度仅在较深层使用 \(h_{\text{base}}\)(见附录 B (https://arxiv.org/html/2606.27449#A2))。所有其他架构超参数在基线和 Prism 变体之间保持一致。
### 3.2 训练结果
为了评估基本的架构影响(原文在此处被截断,但根据上下文,下文应描述训练结果)。相似文章
Asymmetric Attention Heads: 面向Transformer注意力的结构化头级上下文分配
本文介绍了Asymmetric Attention Heads (AAH),一个为transformer中的注意力头分配不同上下文窗口的框架,实验表明语言建模性能得到提升。
@VukRosic99: 长上下文Transformer面临两大瓶颈:二次注意力计算和KV缓存(在1M tokens时可达数百GB)…
MiniCPM-SALA是一款9B参数的混合注意力模型,通过在稀疏注意力和线性注意力之间交替插入(每3个线性层插入1个稀疏层)来克服长上下文Transformer的二次计算和KV缓存瓶颈。在256K tokens下,其推理速度比Qwen3-8B快3.5倍,并能在消费级GPU上支持高达1M tokens。该模型采用经济高效的持续训练方法,训练成本降低约75%。
HydraHead:从头部级功能异质性到专注意力混合
HydraHead 是一种新颖的注意力混合架构,通过在头部层级结合完全注意力和线性注意力,利用可解释性驱动的选择和尺度归一化融合,实现长上下文性能卓越并减少训练开销。
使用稀疏Transformer进行生成建模
OpenAI推出了稀疏Transformer,一种深度神经网络,将注意力机制的复杂度从O(N²)优化到O(N√N),使得能够对长度超过以前30倍的序列进行建模,适用于文本、图像和音频领域。该模型采用稀疏注意力模式和基于检查点的内存优化技术,可以训练深达128层的网络,在多个领域实现了最先进的性能。
Hierarchical Global Attention (HGA)
Hierarchical Global Attention (HGA) 是一种可直接替换预训练长上下文Transformer中密集因果注意力的方法。它采用分层两级路由机制,使得能够对一个小规模路由工作集进行精确注意力计算,从而允许像 Qwen3-30B 这样的模型在单个 RTX 5090 上以64K上下文运行,且质量损失极小。