一、整体架构: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]
用代码表示就是 MiniMindBlock 的 forward:
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_norm 和 k_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。

Comments | NOTHING