mHC 残差流形

mHC_残差流形

1. 先看宏观:mHC 到底改了 Transformer 的哪里?

标准 Transformer 里,每个 Attention 或 FFN/MLP 子层外面都有一个经典残差连接:

\[x_{l+1}=x_l+\mathcal{F}(x_l) \]

mHC 可以理解成:用一套可学习、受约束的“多车道残差路由系统”,替换这个简单加号。

它不改变 Attention 怎么算,也不改变 FFN 怎么算。Attention/FFN 仍然是原来的主干计算 \(\mathcal{F}\)。mHC 改的是主干层外面的残差流拓扑:原来只有一条残差流,现在扩成 \(n\) 条并行残差流,并且让这些流之间可以稳定通信。

2. mHC 作用于谁:残差流从 1 条变成 n 条

传统模型里,残差流是一条维度为 \(C\) 的“单行道”。每一层做的事情大致是:保留旧状态 \(x_l\),再加上 \(\mathcal{F}(x_l)\) 产生的新信息。

mHC 把这条单行道扩成 \(n\) 条并行车道,论文里常见设置是 \(n=4\)

\[x_l \in \mathbb{R}^{B \times T \times n \times C} \]

这里 \(B\) 是 batch,\(T\) 是序列长度,\(C\) 是单条流的 hidden size。直觉上,模型拥有了多条可以持续向后传递的“记忆通道”;但主干计算层依然只接收一条 \(C\) 维输入,所以 mHC 必须负责在多条流和主干层之间做路由。

3. mHC 的四步宏观流程

下面这个流程图把 mHC 放回 LLM 的整体网络拓扑里看:它包裹在每一个 Attention/FFN 子层外面,负责读、算、混、写。

mHC 宏观拓扑流程图

mHC 不改 Attention/FFN 的计算本体,而是改造残差流之间的读、混合、写。

残差流从 1 条扩到 n 条 主干计算仍是 C 维 Hres 受 Birkhoff 约束
输入残差流 xl




Hpre · Read

从 n 条残差流里加权读取,聚合成主干层能吃下的 1 条 C 维输入。

xmain = Hpre · xl




F · Attention / FFN

主干网络照常工作,不需要知道外面有几条残差流。

y = F(xmain)




Hpost · Write

把主干层的新结果按比例广播回 n 条残差流。

Hpost ⊙ y




Hres · 残差内循环混合

行和=1 · 列和=1




Sinkhorn 把原始路由矩阵拉回 Birkhoff polytope,约束车道间交换不整体放大信号。


xl+1 =
Hres · xl
+ Hpost ⊙ F(Hpre · xl)


输出残差流 xl+1




标准 Transformer
一条 C 维残差流:xl → xl + F(xl)
所有信息挤在同一条残差通道里,模型能力主要靠加深或加宽主干网络扩展。
mHC Transformer
残差状态变宽,主干计算宽度保持 C;额外成本集中在很小的路由矩阵和门控上。
输入状态 x: [B,T,n,C]
Read 门控 H_pre: [B,T,n,1]
主干输入 sum_n(H_pre*x): [B,T,C]
Mix 矩阵 H_res: [B,T,n,n]
输出状态 x_next: [B,T,n,C]

3.1 Read / Pre-routing:用 \(H_{pre}\) 聚合输入

主干层 \(\mathcal{F}\) 只吃一条 \(C\) 维输入,但 mHC 当前有 \(n\) 条残差流。于是它用 \(H_{pre}\) 做加权读取,把多条流聚合成一条主干输入:

\[x_{\text{main}} = H_{pre} \cdot x_l \]

可以把 \(H_{pre}\) 理解成“读门控”:每一层、每个 token 都可以动态决定从哪几条残差流里多读一点。

3.2 Compute:Attention / FFN 正常计算

主干层拿到聚合后的 \(x_{\text{main}}\),照常执行 Attention 或 FFN:

\[y = \mathcal{F}(x_{\text{main}}) \]

这一步是 mHC 很重要的工程取舍:它没有把 Attention/FFN 的宽度也扩成 \(nC\),因此主要 FLOPs 仍然接近原模型。

3.3 Mix:用 \(H_{res}\) 让残差流互相通信

在主干层计算的同时,外面的 \(n\) 条残差流也会通过 \(H_{res}\) 互相混合:

\[x_{\text{res}} = H_{res} \cdot x_l \]

这里的 \(H_{res}\) 是 mHC 的数学核心。它被约束成双随机矩阵:

  • 元素非负;
  • 每一行和为 1;
  • 每一列和为 1。

也就是说,\(H_{res}\) 位于 Birkhoff Polytope 内。这样做的目的不是为了“好看”,而是为了让流与流之间的信息交换保持非扩张性质:信息可以重新分配,但整体不容易被层层放大到失控。

3.4 Write / Post-routing:用 \(H_{post}\) 分发主干输出

主干层输出 \(y\) 后,mHC 用 \(H_{post}\) 把这份新信息广播并写回到 \(n\) 条残差流:

\[x_{l+1} = H_{res} \cdot x_l + H_{post} \odot \mathcal{F}(H_{pre} \cdot x_l) \]

所以 mHC 的一层不是“旧状态 + 新状态”这么简单,而是:

  • 从多条残差流读取;
  • 用原来的 Attention/FFN 计算;
  • 让残差流之间稳定混合;
  • 把新信息按比例写回多条残差流。

4. 为什么要引入这套机制?

传统扩展模型能力通常有两条路:

  • 做深:增加层数;
  • 做宽:增大 hidden size \(C\)

但这两条路都会显著增加计算量,尤其是把主干层直接做宽时,矩阵乘法 FLOPs 会快速上升。Hyper-Connections 的思路是:尽量保持主干计算层不变,只拓宽残差状态本身。

这带来一个很诱人的收益:模型拥有更多并行残差记忆通道,但 Attention/FFN 的主要计算宽度仍然是 \(C\)

问题是,如果多条残差流随便混合,深层堆叠后很容易破坏残差网络原本依赖的恒等映射性质,导致训练不稳定、梯度爆炸或数值漂移。mHC 的核心贡献就是给这套多车道残差系统加上约束:

  • 数学交规:用 Birkhoff Polytope / 双随机矩阵约束 \(H_{res}\),保证残差混合更稳定;
  • 工程优化:通过算子融合和重计算降低多条残差流带来的显存带宽与显存占用压力;
  • 软启动:用很小的 \(\alpha\) 初始化路由强度,让模型初始行为接近标准 ResNet/Transformer 残差连接。

一句话总结:

mHC 不是改变模型“怎么思考”,而是改变模型内部“记忆如何并行保存、通信和传递”。

5. 与 Birkhoff Polytope 的联系

Birkhoff Polytope 是所有双随机矩阵构成的凸多面体,它的顶点是所有置换矩阵。

在 mHC 中,把 \(H_{res}\) 约束到 Birkhoff Polytope 内,有一个非常直接的残差流解释:

  • 置换矩阵只重排残差流,不改变整体大小;
  • 双随机矩阵是置换矩阵的凸组合,可以看成“软重排”;
  • 因此 \(H_{res}\) 可以让多条残差流互相通信,同时不轻易放大信号。

Sinkhorn-Knopp 算法就是把一个普通的可学习矩阵反复做行归一化、列归一化,最终拉回到双随机矩阵集合附近:

H = torch.exp(M / tau)
for _ in range(n_iters):
    H = H / (H.sum(dim=-1, keepdim=True) + 1e-8)
    H = H / (H.sum(dim=-2, keepdim=True) + 1e-8)

这就是前面流程图里 \(H_{res}\) 那个“稳定混合矩阵”的来源。

6. 代码形状对照

如果把上面的宏观流程翻译成 PyTorch 张量形状,大致是:

# x: [B, T, n, C]

H_pre = sigmoid(alpha_pre * proj_pre(x_global))      # [B, T, n, 1]
x_main = (x * H_pre).sum(dim=2)                      # [B, T, C]

y = layer_func(x_main)                               # [B, T, C]

H_res_raw = proj_res(x_global).view(B, T, n, n)
H_res = sinkhorn_knopp(alpha_res * H_res_raw)         # [B, T, n, n]
x_res = torch.matmul(H_res, x)                        # [B, T, n, C]

H_post = 2 * sigmoid(alpha_post * proj_post(x_global))# [B, T, n, 1]
x_next = x_res + H_post * y.unsqueeze(2)              # [B, T, n, C]

这一段代码对应的就是:

\[x_{l+1}=H_{res}x_l+H_{post}\odot \mathcal{F}(H_{pre}x_l) \]

7. 几何可视化

上面的流程图解决的是“mHC 在整台 Transformer 发动机里装在哪里”。下面两个图则继续深入那个齿轮本身:为什么 \(H_{res}\) 要被约束到 Birkhoff Polytope,以及这个几何对象长什么样。

Birkhoff Polytope 交互演示

① 什么是双随机矩阵
满足三条约束:每行和 = 1每列和 = 1所有元素 >= 0。编辑下面 3x3 矩阵,观察约束是否成立。



② 2x2 情形:一条线段

2x2 双随机矩阵只有 1 个自由度:B(t) = [[t, 1-t], [1-t, t]],其中 t 属于 [0,1]。




t =

0.50





③ 3x3 情形:6 个置换矩阵顶点

3x3 的所有置换矩阵一共 6 个,它们是 Birkhoff 多面体的顶点。点击任意顶点,查看对应矩阵和投影位置。




当前顶点:P1









④ 凸组合:任意双随机矩阵都可分解

Birkhoff-von Neumann 定理:任意双随机矩阵可写为置换矩阵的非负加权和,且权重和为 1。













⑤ Sinkhorn:把任意正矩阵拉回双随机集合

交替执行行归一化与列归一化。点击按钮观察偏差如何收敛。














H = exp(M / tau)\nfor k in range(K):\n H = H / row_sum(H)\n H = H / col_sum(H)






⑥ mHC 中的作用:稳定混合残差流

在 mHC 中,H_res 被约束为双随机矩阵,可以在不显著放大信号的前提下完成多流通信。







矩阵约束作用
H_res双随机流与流之间的软重排,保证稳定混合
H_pre[0,1] 门控从多流读入主干计算
H_post[0,2] 门控把主干输出分发回多流


x_{l+1} = H_res x_l + H_post ⊙ F(H_pre x_l)

你可以把它理解为:先混流,再把主干计算结果按比例写回,整体比无约束混合更稳。


3D Birkhoff 投影演示

3D 视角看懂 Birkhoff 多面体

把 3x3 双随机矩阵看作 6 个置换矩阵的凸组合:顶点是纯置换矩阵,中间蓝点是当前混合结果。



拖动滑块调整比例,按住画布旋转 3D 视角

P1

P2

P3

P4

P5

P6

混合矩阵


完整代码练习

import torch
import torch.nn as nn

# ==========================================
# 1. 前置依赖 (沿用之前的核心逻辑)
# ==========================================
def sinkhorn_knopp(M, n_iters=20, tau=1.0):
    H = torch.exp(M / tau)
    for _ in range(n_iters):
        H = H / (H.sum(dim=-1, keepdim=True) + 1e-8)
        H = H / (H.sum(dim=-2, keepdim=True) + 1e-8)
    return H

class mHCBlock(nn.Module):
    def __init__(self, dim, n_streams=4):
        super().__init__()
        self.dim = dim
        self.n = n_streams
        total_dim = n_streams * dim
        
        self.norm = nn.RMSNorm(total_dim)
        self.proj_pre  = nn.Linear(total_dim, n_streams)
        self.proj_post = nn.Linear(total_dim, n_streams)
        self.proj_res  = nn.Linear(total_dim, n_streams * n_streams)
        
        self.alpha_pre  = nn.Parameter(torch.tensor(0.01))
        self.alpha_post = nn.Parameter(torch.tensor(0.01))
        self.alpha_res  = nn.Parameter(torch.tensor(0.01))

    def forward(self, x, layer_func):
        B, L, N, C = x.shape
        x_global = x.view(B, L, -1)
        x_norm = self.norm(x_global)
        
        # 计算路由权重
        H_pre = torch.sigmoid(self.alpha_pre * self.proj_pre(x_norm)).view(B, L, N, 1)
        H_post = 2.0 * torch.sigmoid(self.alpha_post * self.proj_post(x_norm)).view(B, L, N, 1)
        H_res = sinkhorn_knopp(self.alpha_res * self.proj_res(x_norm).view(B, L, N, N))

        # A. 聚合
        x_main_in = (x * H_pre).sum(dim=2) 
        # B. 主干计算
        x_main_out = layer_func(x_main_in)
        # C. 残差互通
        x_res_mixed = torch.matmul(H_res, x)
        # D. 广播分发
        x_next = x_res_mixed + H_post * x_main_out.unsqueeze(2)
        
        return x_next

# ==========================================
# 2. 核心处理厂:标准的前馈神经网络 (FFN)
# ==========================================
class FeedForward(nn.Module):
    """标准的 MLP 处理模块,它完全不需要知道外面有 mHC 的存在"""
    def __init__(self, dim, hidden_dim):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(dim, hidden_dim),
            nn.GELU(),
            nn.Linear(hidden_dim, dim)
        )
        
    def forward(self, x):
        return self.net(x)

# ==========================================
# 3. 宏观架构:基于 mHC 的分类器
# ==========================================
class mHC_SequenceClassifier(nn.Module):
    def __init__(self, vocab_size, dim, n_streams, num_layers, out_classes):
        super().__init__()
        self.dim = dim
        self.n = n_streams
        
        # 1. 输入层:将词元映射为向量 [B, L, C]
        self.embedding = nn.Embedding(vocab_size, dim)
        
        # 2. 拓宽车道:把 1 条 C 维的车道,强行拆解成 n 条 C 维的车道
        self.expand = nn.Linear(dim, n_streams * dim)
        
        # 3. 堆叠 mHC 层和 FFN 层
        self.mhc_blocks = nn.ModuleList([mHCBlock(dim, n_streams) for _ in range(num_layers)])
        self.ffn_layers = nn.ModuleList([FeedForward(dim, dim * 4) for _ in range(num_layers)])
        
        # 4. 压缩车道:把 n 条车道的信息压平,合并回 1 条 C 维车道
        self.collapse = nn.Linear(n_streams * dim, dim)
        
        # 5. 输出层:简单的池化 + 分类头
        self.head = nn.Linear(dim, out_classes)

    def forward(self, x):
        B, L = x.shape
        
        # [B, L] -> [B, L, C]
        x = self.embedding(x)
        
        # [B, L, C] -> [B, L, N*C] -> [B, L, N, C]
        # 此时数据正式进入 4 条并行的残差流
        x = self.expand(x).view(B, L, self.n, self.dim)
        
        # 逐层穿过 mHC 模块
        for mhc, ffn in zip(self.mhc_blocks, self.ffn_layers):
            # 将 ffn 作为 callable 函数传给 mhc
            x = mhc(x, ffn)
            
        # 离开多数据流区域,将 [B, L, N, C] 压平为 [B, L, N*C]
        x = x.view(B, L, -1)
        
        # 合并回单主干道 [B, L, C]
        x = self.collapse(x)
        
        # 序列池化:取所有 Token 的平均特征 [B, C]
        x_pooled = x.mean(dim=1)
        
        # 输出预测结果 [B, out_classes]
        logits = self.head(x_pooled)
        return logits

# ==========================================
# 4. 测试与验证
# ==========================================
if __name__ == "__main__":
    # 超参数设置
    BATCH_SIZE = 8
    SEQ_LEN = 128
    VOCAB_SIZE = 1000
    DIM = 256
    N_STREAMS = 4
    NUM_LAYERS = 6
    NUM_CLASSES = 10

    # 实例化模型
    model = mHC_SequenceClassifier(
        vocab_size=VOCAB_SIZE, 
        dim=DIM, 
        n_streams=N_STREAMS, 
        num_layers=NUM_LAYERS, 
        out_classes=NUM_CLASSES
    )
    
    # 将模型扔到 GPU 上(如果你在 Ubuntu 环境中执行,可以开启)
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    model = model.to(device)
    
    # 构造假输入数据
    dummy_input = torch.randint(0, VOCAB_SIZE, (BATCH_SIZE, SEQ_LEN)).to(device)
    
    # 前向传播
    print("🚀 开始前向传播...")
    outputs = model(dummy_input)
    print(f"✅ 输出张量维度: {outputs.shape} (预期: [{BATCH_SIZE}, {NUM_CLASSES}])")
    
    # 模拟一次反向传播
    print("\n🚀 模拟计算 Loss 并反向传播...")
    loss = outputs.sum()
    loss.backward()
    print("✅ 反向传播完成!")
    
    # 检查梯度是否正常(通过 Sinkhorn 约束,梯度应该非常稳定)
    grad_norm = model.mhc_blocks[-1].proj_res.weight.grad.norm().item()
    print(f"📊 最后一层 mHC proj_res 的梯度范数: {grad_norm:.4f}")
	
posted @ 2026-07-18 05:17  然皓  阅读(20)  评论(0)    收藏  举报