Hyper-Connections:残差连接的可学习替代方案

项目 内容
论文 Hyper-Connections
会议 ICLR 2025
作者 Defa Zhu, Hongzhi Huang, Zihao Huang, Yutao Zeng, Yunyao Mao, Banggu Wu, Qiyang Min, Xun Zhou
机构 Seed-Foundation-Model Team, ByteDance
arXiv 2409.19606

问题背景

残差连接(Residual Connections)是深度网络训练的基石,但存在根本性的跷跷板效应(seesaw effect)

  • Pre-Norm:在残差块前做归一化,解决梯度消失,但导致表示坍塌——深层隐状态高度相似,层间贡献递减
  • Post-Norm:在残差块后做归一化,缓解表示坍塌,但重新引入梯度消失

两种方案的连接强度都是预定义的、不可训练的。核心问题:网络能否自主学习最优的连接强度?

方法

核心思想

将单一残差流扩展为 nn 个并行流(expansion rate = nn),通过可学习的残差矩阵动态混合信息。

形式化定义

隐状态 hk1Rd\mathbf{h}^{k-1} \in \mathbb{R}^d 被复制 nn 次形成 hyper hidden matrix H=(h1,h2,,hn)Rn×d\mathbf{H} = (\mathbf{h}_1, \mathbf{h}_2, \ldots, \mathbf{h}_n)^\top \in \mathbb{R}^{n \times d}

Hyper-Connection 矩阵 HCR(n+1)×(n+1)\mathcal{HC} \in \mathbb{R}^{(n+1) \times (n+1)}

HC=(01×1BAmAr)\mathcal{HC} = \begin{pmatrix} \mathbf{0}_{1\times1} & \mathbf{B} \\ \mathbf{A_m} & \mathbf{A_r} \end{pmatrix}

输出计算:H^=BT(HAm)+ArH\hat{\mathbf{H}} = \mathbf{B}^\top \mathcal{T}(\mathbf{H}^\top \mathbf{A_m})^\top + \mathbf{A_r}^\top \mathbf{H}

其中:

  • Am\mathbf{A_m}:加权求和输入 H\mathbf{H} 得到当前层输入 h0\mathbf{h}_0深度连接
  • Ar\mathbf{A_r}:将 H\mathbf{H} 映射为 hyper hidden matrix H\mathbf{H}'宽度连接
  • B\mathbf{B}:将层输出分配回各流

可分解为两个子连接

连接类型 作用 矩阵
Depth-Connections 跨深度的加权残差,控制各层对当前层的贡献 DC=(B;diag(Ar))R2×n\mathcal{DC} = (\mathbf{B}; \text{diag}(\mathbf{A_r})) \in \mathbb{R}^{2 \times n}
Width-Connections 同层内多流间的信息交换 WC=(Am,Ar)Rn×(n+1)\mathcal{WC} = (\mathbf{A_m}, \mathbf{A_r}) \in \mathbb{R}^{n \times (n+1)}

动态 Hyper-Connections (DHC)

矩阵元素可依赖输入 H\mathbf{H} 动态生成:

B(H)=sβtanh(HˉWβ)+B\mathcal{B}(\mathbf{H}) = s_\beta \circ \tanh(\bar{\mathbf{H}} \mathbf{W}_\beta)^\top + \mathbf{B}
Am(H)=sαtanh(HˉWm)+Am\mathcal{A}_m(\mathbf{H}) = s_\alpha \circ \tanh(\bar{\mathbf{H}} \mathbf{W}_m) + \mathbf{A}_m
Ar(H)=sαtanh(HˉWr)+Ar\mathcal{A}_r(\mathbf{H}) = s_\alpha \circ \tanh(\bar{\mathbf{H}} \mathbf{W}_r) + \mathbf{A}_r

其中 Hˉ=norm(H)\bar{\mathbf{H}} = \text{norm}(\mathbf{H})sβ,sαs_\beta, s_\alpha 为可学习缩放因子。

初始化策略

初始化等价于 Pre-Norm 残差连接:

(01×1BkAmkArk)=(01×111×nekmodnen×n)\begin{pmatrix} \mathbf{0}_{1\times1} & \mathbf{B}^k \\ \mathbf{A_m}^k & \mathbf{A_r}^k \end{pmatrix} = \begin{pmatrix} \mathbf{0}_{1\times1} & \mathbf{1}_{1\times n} \\ \mathbf{e}_{k \bmod n} & \mathbf{e}_{n \times n} \end{pmatrix}

残差连接是 HC 的特例

方法 对应 HC 矩阵 (n=1n=1)
Pre-Norm (0111)\begin{pmatrix} 0 & 1 \\ 1 & 1 \end{pmatrix}
Post-Norm (01σi2+σo2+2σio11σi2+σo2+2σio)\begin{pmatrix} 0 & \frac{1}{\sqrt{\sigma_i^2+\sigma_o^2+2\sigma_{io}}} \\ 1 & \frac{1}{\sqrt{\sigma_i^2+\sigma_o^2+2\sigma_{io}}} \end{pmatrix}

Sequential-Parallel Duality

HC 可以学习出介于顺序和并行之间的层排列:

  • 顺序排列:HC=(011110001)\mathcal{HC} = \begin{pmatrix} 0 & 1 & 1 \\ 1 & 1 & 0 \\ 0 & 0 & 1 \end{pmatrix}(退化为标准残差)
  • 并行排列:奇偶层使用不同矩阵,等价于 Parallel Transformer Block

实验结果

消融实验(OLMo-1B, 500B tokens)

方法 V2 Eval Loss V2 PPL V3 Eval Loss V3 PPL 下游平均 Acc.
OLMo-1B (baseline) 2.811 18.023 2.544 14.229 62.5
DHC×1 2.819 18.125 2.556 14.418 62.3
DHC×2 2.802 17.950 2.534 14.114 63.0
DHC×4 2.781 17.509 2.516 13.826 63.8
DHC×8 2.778 17.445 2.516 13.843 62.8
DHC×4 W/O tanh 2.779 17.451 2.516 13.844 64.4
DHC×8 W/O tanh 2.777 17.425 2.514 13.819 63.8

关键发现:n=1n=1 时性能不如 baseline(跷跷板效应仍在),n2n \geq 2 时显著超越。

7B Dense 模型

方法 Params (B) FLOPs (G) V2 Loss V2 PPL V3 Loss V3 PPL 下游 Avg.
OLMo-7B 6.9 13.36 2.581 14.316 2.322 11.324 70.1
OLMo-7B-DHC×4 6.9 13.38 2.559 14.023 2.304 11.120 71.0

参数量和计算量几乎不变,但 loss 和下游任务均有提升。

MoE 模型(OLMoE-1B-7B, 500B tokens)

方法 MMLU Var HellaSwag ARC-C ARC-E PIQA WinoGrande BoolQ
OLMoE-1B-7B 38.5 69.5 41.8 72.8 77.6 64.4 65.4
OLMoE-1B-7B-DHC×4 39.7 70.2 47.8 76.7 78.2 64.6 68.5

MoE 模型收益更大:ARC-Challenge 提升 6 分,收敛速度快 1.8 倍。

与相关方法对比

方法 V2 Loss V2 PPL V3 Loss V3 PPL Avg. Acc.
OLMo-1B 2.811 18.023 2.544 14.229 62.5
ResiDual 2.825 18.375 2.551 14.346 62.0
Altup×2 2.827 18.268 2.558 14.454 62.4
DHC×2 2.802 17.950 2.534 14.114 63.0

ResiDual 和 Altup 在训练初期有增益但最终被 baseline 超越,HC 则持续领先。

核心洞察

可视化分析(连接矩阵)

通过展开 HC 为 dense connection weight matrix,观察到:

  1. Λ 形连接模式:长程衰减(Post-Norm 风格)+ 底层频繁访问(Pre-Norm 风格),HC 自动学出了两者的混合
  2. 输入 embedding 被消除:word embedding 对大多数层贡献极小(仅最后一层除外),保留 embedding 分量有害于 next token prediction
  3. 并行 Transformer Block 自发涌现:相邻层间出现锯齿形模式(如 layer 11 对 layer 12 贡献极小),等价于 attention+FFN 并行执行
  4. Attention 层长程连接更少:attention 层底部几乎无长程贡献,FFN 层输出幅度显著更大,类似 two-hop residual 设计

表示坍塌缓解

Figure 3 显示:Pre-Norm 模型相邻层余弦相似度中位数 > 0.6(严重坍塌),HC 模型降至 ~0.2-0.4,层间差异性显著增大。

设计要点总结

设计选择 结论
Expansion rate n=4n=4 为最佳性价比,n=8n=8 边际收益小
Static vs Dynamic DHC 优于 SHC,尤其在 n=4n=4 时差距明显
tanh 激活 有 tanh 时 PPL 更优,无 tanh 时下游 Acc 略高
WC\mathcal{WC} 可训练性 关键——不训练 WC\mathcal{WC} 导致 V2 loss 增加 0.021
B\mathbf{B} 可训练性 重要但影响略小于 WC\mathcal{WC}
计算开销 参数和 FLOPs 增加可忽略(7B 模型仅 +0.02G FLOPs)
AI Assistant
Selected Text