MiniMind的block

发布于 15 天前  49 次阅读


一、整体架构:Pre-Norm 的 Decoder Block

MiniMind 的 Block 采用 Pre-Norm 架构(先归一化再计算),这是当前主流 LLM(LLaMA、Qwen3 等)的标准设计。一个 Block 的数据流如下:

输入 hidden_states [B, S, 768]
        │
        ├──→ RMSNorm (input_layernorm) ──→ GQA Attention ──┐
        │                                                    │
        └──────────────────── 残差相加 ←─────────────────────┘
                                │
        ├──→ RMSNorm (post_attention_layernorm) ──→ SwiGLU FFN ──┐
        │                                                         │
        └──────────────────────── 残差相加 ←───────────────────────┘
                                │
                        输出 hidden_states [B, S, 768]

用代码表示就是 MiniMindBlockforward

def forward(self, hidden_states, position_embeddings, past_key_value=None, use_cache=False, attention_mask=None):
    residual = hidden_states                                    # ① 保存残差
    hidden_states, present_key_value = self.self_attn(
        self.input_layernorm(hidden_states), position_embeddings,  # ② Pre-Norm + Attention
        past_key_value, use_cache, attention_mask
    )
    hidden_states += residual                                  # ③ 第一次残差连接

    hidden_states = hidden_states + self.mlp(
        self.post_attention_layernorm(hidden_states)            # ④ Pre-Norm + FFN
    )                                                          # ⑤ 第二次残差连接
    return hidden_states, present_key_value

二、逐组件详解

1. RMSNorm(Root Mean Square Normalization)

Block 里有两个 RMSNorm,分别位于 Attention 前和 FFN 前。

class RMSNorm(torch.nn.Module):
    def __init__(self, dim: int, eps: float = 1e-5):
        super().__init__()
        self.eps = eps
        self.weight = nn.Parameter(torch.ones(dim))   # 只有 γ,没有 β

    def norm(self, x):
        return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)

    def forward(self, x):
        return (self.weight * self.norm(x.float())).type_as(x)

为什么用 RMSNorm 而不是 LayerNorm?

  • LayerNorm 计算:$\text{LN}(x) = \gamma \cdot \frac{x - \mu}{\sqrt{\sigma^2 + \epsilon}} + \beta$
  • RMSNorm 计算:$\text{RMSNorm}(x) = \gamma \cdot \frac{x}{\sqrt{\text{RMS}(x)^2 + \epsilon}}$,其中 $\text{RMS}(x) = \sqrt{\frac{1}{n}\sum x_i^2}$

RMSNorm 去掉了均值偏移(减去 $\mu$)和偏置项($\beta$),在保持训练稳定性的同时减少了计算量。实验表明,在 Transformer 中去掉均值中心化对性能影响不大,但能加速训练。

参数:每个 RMSNorm 只有 hidden_size=768 个可学习参数($\gamma$),两个合计 1,536 参数。


2. GQA Attention(Grouped Query Attention)

这是 Block 的核心计算单元,我们之前详细讨论过。在 Block 中,Attention 的输入是经过 input_layernorm 后的 hidden_states。

class Attention(nn.Module):
    def __init__(self, config: MiniMindConfig):
        self.q_proj = nn.Linear(768, 768, bias=False)      # Q: 768 → 768 (8 heads × 96 dim)
        self.k_proj = nn.Linear(768, 384, bias=False)      # K: 768 → 384 (4 heads × 96 dim)
        self.v_proj = nn.Linear(768, 384, bias=False)      # V: 768 → 384 (4 heads × 96 dim)
        self.o_proj = nn.Linear(768, 768, bias=False)      # O: 768 → 768
        self.q_norm = RMSNorm(96, eps=1e-6)                # Query 头维度归一化
        self.k_norm = RMSNorm(96, eps=1e-6)                # Key 头维度归一化
关键配置 参数 含义
num_attention_heads 8 Query 头数
num_key_value_heads 4 KV 头数(GQA 分组)
head_dim 96 每个头的维度
n_rep 2 KV 头复制倍数(8/4=2)

数据流维度变化(以 batch_size=B, seq_len=S 为例):

输入 x:              [B, S, 768]
        │
        ├──→ q_proj ──→ [B, S, 768] ──view──→ [B, S, 8, 96]  (8个Query头)
        ├──→ k_proj ──→ [B, S, 384] ──view──→ [B, S, 4, 96]  (4个KV头)
        └──→ v_proj ──→ [B, S, 384] ──view──→ [B, S, 4, 96]
        │
        ├──→ q_norm/k_norm (头维度RMSNorm)
        ├──→ RoPE 位置编码 (apply_rotary_pos_emb)
        ├──→ repeat_kv: K/V [B, S, 4, 96] ──→ [B, S, 8, 96] (复制2倍)
        │
        └──→ Attention 计算 ──→ [B, S, 8, 96] ──reshape──→ [B, S, 768]
                    │
                    └──→ o_proj ──→ [B, S, 768]

Attention 内部还有一个关键细节q_normk_norm。这是 Qwen3 引入的头维度归一化,在 RoPE 之前对 Q/K 的每个头做 RMSNorm,能提升训练稳定性。

参数量

  • W_Q: 768 × 768 = 589,824
  • W_K: 768 × 384 = 294,912
  • W_V: 768 × 384 = 294,912
  • W_O: 768 × 768 = 589,824
  • Attention 合计:1,769,472(≈1.77M)

3. SwiGLU FFN(Feed-Forward Network)

Block 中的第二个计算单元,在 Attention 之后。

class FeedForward(nn.Module):
    def __init__(self, config: MiniMindConfig, intermediate_size: int = None):
        intermediate_size = intermediate_size or config.intermediate_size  # 2432
        self.gate_proj = nn.Linear(768, 2432, bias=False)   # 门控投影
        self.up_proj   = nn.Linear(768, 2432, bias=False)   # 升维投影
        self.down_proj = nn.Linear(2432, 768, bias=False)   # 降维投影
        self.act_fn = ACT2FN['silu']                        # SiLU 激活

    def forward(self, x):
        return self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x))

intermediate_size 的计算

intermediate_size = math.ceil(hidden_size * math.pi / 64) * 64
                  = math.ceil(768 * 3.14159 / 64) * 64
                  = math.ceil(37.699) * 64
                  = 38 * 64 = 2432

这里用 $\pi$ 来确定 FFN 中间维度是一个有趣的设计选择,使得 intermediate_size ≈ 3.14 × hidden_size

SwiGLU 的计算逻辑: $$\text{FFN}(x) = W{\text{down}} \cdot (\text{SiLU}(W{\text{gate}} \cdot x) \odot (W_{\text{up}} \cdot x))$$

其中 $\odot$ 是逐元素乘法(Hadamard 积)。SiLU 作为门控信号,决定哪些信息通过。

参数量

  • gate_proj: 768 × 2432 = 1,867,776
  • up_proj: 768 × 2432 = 1,867,776
  • down_proj: 2432 × 768 = 1,867,776
  • FFN 合计:5,603,328(≈5.60M)

4. 残差连接(Residual Connection)

Block 中有两处残差连接:

# 第一处:Attention 之后
hidden_states += residual

# 第二处:FFN 之后  
hidden_states = hidden_states + self.mlp(...)

残差连接的核心作用是缓解梯度消失,让梯度可以直接回传到浅层。在 Pre-Norm 架构中,残差连接绕过归一化层,保证了信息流的畅通。


三、一个 Block 的完整数据流(带维度)

假设输入 hidden_states 的 shape 为 [B, S, 768]

步骤 操作 输入维度 输出维度 备注
1 input_layernorm [B, S, 768] [B, S, 768] RMSNorm
2 self_attn [B, S, 768] [B, S, 768] GQA + RoPE
3 残差相加 [B, S, 768] + [B, S, 768] [B, S, 768] 第一处残差
4 post_attention_layernorm [B, S, 768] [B, S, 768] RMSNorm
5 mlp (SwiGLU) [B, S, 768] [B, S, 768] FFN
6 残差相加 [B, S, 768] + [B, S, 768] [B, S, 768] 第二处残差

可以看到,Block 的输入输出维度完全一致([B, S, 768]),这使得多层 Block 可以堆叠串联


四、参数量汇总

以 MiniMind-3 的配置(hidden_size=768, num_hidden_layers=8)计算:

组件 参数量 占比
Attention 1,769,472 24.0%
FFN (SwiGLU) 5,603,328 76.0%
RMSNorm (×2) 1,536 ~0%
单 Block 合计 7,374,336 100%
8 层 Block 总计 58,994,688 ≈59M

关键发现

  • FFN 的参数量是 Attention 的 3.17 倍,这是现代 Transformer 的普遍特征(FFN 负责存储知识,Attention 负责路由)。
  • 所有 Linear 层都 没有 bias,节省了约 10% 的参数。
  • GQA 让 K/V 投影参数量只有 Q 投影的一半(384 vs 768)。

五、与原始 Transformer 的对比

特性 原始 Transformer (2017) MiniMind Block
归一化位置 Post-Norm(计算后归一化) Pre-Norm(计算前归一化)
归一化类型 LayerNorm RMSNorm
注意力 Multi-Head Attention (MHA) GQA (8Q/4KV)
位置编码 正弦/余弦绝对位置编码 RoPE 旋转位置编码
FFN 激活 ReLU SwiGLU (SiLU + Gating)
FFN 结构 两层 Linear 三层 Linear (gate/up/down)
残差连接

六、Block 在整体模型中的位置

class MiniMindModel(nn.Module):
    def __init__(self, config: MiniMindConfig):
        self.embed_tokens = nn.Embedding(6400, 768)           # 词嵌入
        self.layers = nn.ModuleList([
            MiniMindBlock(l, config) for l in range(8)        # 8个Block堆叠
        ])
        self.norm = RMSNorm(768, eps=1e-6)                    # 最终RMSNorm

    def forward(self, input_ids, ...):
        hidden_states = self.embed_tokens(input_ids)           # [B, S, 768]
        for layer in self.layers:
            hidden_states, _ = layer(hidden_states, ...)       # 逐层通过8个Block
        hidden_states = self.norm(hidden_states)               # 最终归一化
        return hidden_states

8 个 Block 串联后,输出再经过最终的 RMSNorm,然后送给 lm_head 做 Next Token Prediction。


一沙一世界,一花一天堂。君掌盛无边,刹那成永恒。