Stephen 技术博客

AI和数据平台的工程实践

← 返回文章列表

Makemore Part 3 总结:激活、初始化、BatchNorm

对应 Karpathy "Building makemore Part 3: Activations & Gradients, BatchNorm"
学习日:Day 13–15

一、为什么从 Part 2 升级到 Part 3

Part 2 的局限性:

维度Part 2 (MLP)Part 3 (诊断 + 修复)
网络深度1 层 hidden可以堆 6+ 层
初始化默认 randn,巧合能跑系统的 Kaiming 初始化
激活值监控看 loss 就行看每层激活值/梯度直方图
训练稳定性浅层不会爆炸加 BatchNorm 才能堆深
调试思路"loss 不降 → 调 lr""看诊断图 → 定位问题层"
核心洞察:之前两节你的网络"碰巧"能训练,是因为它太浅了。Part 3 揭示了为什么深网络会崩溃,以及现代深度学习的两大支柱(好的初始化 + Normalization)是怎么解决这个问题的。

从这一节开始你的段位提升:不再只会调超参,而是能看诊断图判断网络健康

二、四大主题地图

问题 1: 初始 loss 异常大 (27 而不是 3.29)
   └─ 修复: W2 *= 0.01,让初始 logits 接近 0

问题 2: tanh 饱和,神经元死亡
   └─ 修复: W1 缩小,让 tanh 输入落在非饱和区

问题 3: 凭直觉缩小不靠谱
   └─ 修复: Kaiming 初始化,公式化解决方案
                              ↓
                       但只能保证"训练初期"稳定
                              ↓
问题 4: 训练几步后权重变化,方差又失控
   └─ 修复: BatchNorm,每层强制归一化
                              ↓
                     现代深度学习能堆 100+ 层的关键
                              ↓
问题 5: 怎么知道网络是否健康?
   └─ 工具: 激活值/梯度直方图、update/param ratio

三、关键超参数

参数取值含义
block_size3上下文窗口
embedding_dim10嵌入维度
hidden_size200隐藏层神经元数
vocab_size27字母表大小
batch_size32mini-batch 大小
gain (tanh)5/3Kaiming 的激活修正
epsilon (BN)1e-5BatchNorm 防止除 0
momentum (BN)0.001running mean/var 的更新率

四、关键代码

1. 初始化修复:从默认到 Kaiming

默认(坏)

W1 = torch.randn((30, 200))     # 方差 1,会让 tanh 饱和
W2 = torch.randn((200, 27))     # 方差 1,初始 loss 巨大
b2 = torch.randn(27)

Kaiming 初始化(好)

W1 = torch.randn((30, 200)) * (5/3) / 30**0.5    # tanh: gain=5/3
W2 = torch.randn((200, 27)) * 0.01               # 输出层缩小让 loss≈3.29
b2 = torch.zeros(27)                              # bias 设 0

手算合理初始 loss

# 27 类均匀预测的交叉熵
expected_loss = -math.log(1/27)   # ≈ 3.29

2. BatchNorm 完整实现

class BatchNorm1d:
    def __init__(self, dim, eps=1e-5, momentum=0.001):
        self.eps = eps
        self.momentum = momentum
        self.training = True
        # 可学参数
        self.gamma = torch.ones(dim)
        self.beta = torch.zeros(dim)
        # buffer (推理用)
        self.running_mean = torch.zeros(dim)
        self.running_var  = torch.ones(dim)

    def __call__(self, x):
        if self.training:
            xmean = x.mean(0, keepdim=True)
            xvar  = x.var(0, keepdim=True)
        else:
            xmean = self.running_mean
            xvar  = self.running_var

        xhat = (x - xmean) / torch.sqrt(xvar + self.eps)
        out  = self.gamma * xhat + self.beta

        if self.training:
            with torch.no_grad():
                self.running_mean = (1 - self.momentum) * self.running_mean + self.momentum * xmean
                self.running_var  = (1 - self.momentum) * self.running_var  + self.momentum * xvar

        return out

    def parameters(self):
        return [self.gamma, self.beta]

用法(注意 Linear 不要 bias)

layers = [
    Linear(30, 200, bias=False), BatchNorm1d(200), Tanh(),
    Linear(200, 200, bias=False), BatchNorm1d(200), Tanh(),
    Linear(200, 27)
]

3. PyTorch 风格的模块化

class Linear:
    def __init__(self, fan_in, fan_out, bias=True):
        self.weight = torch.randn((fan_in, fan_out)) / fan_in**0.5    # Kaiming
        self.bias = torch.zeros(fan_out) if bias else None

    def __call__(self, x):
        self.out = x @ self.weight
        if self.bias is not None:
            self.out += self.bias
        return self.out

    def parameters(self):
        return [self.weight] + ([] if self.bias is None else [self.bias])


class Tanh:
    def __call__(self, x):
        self.out = torch.tanh(x)
        return self.out

    def parameters(self):
        return []

五、四种诊断图:代码与输出

按 Karpathy notebook 风格展示:每个诊断都是一段能直接运行的诊断代码,下方紧跟它的图形输出。图用 Kaiming 初始化的 6 层 tanh MLP(dim=30 → 100×5 → 27)生成,呈现"健康"状态。

诊断 1:激活分布图(activation distribution)

1plt.figure(figsize=(20, 4)) 2legends = [] 3for i, layer in enumerate(layers[:-1]): # 遍历所有 tanh 层 4 if isinstance(layer, Tanh): 5 t = layer.out 6 print(f'layer {i} (Tanh): mean {t.mean():+.2f}, std {t.std():.2f}, ' 7 f'saturated: {(t.abs() > 0.97).float().mean()*100:.2f}%') 8 hy, hx = torch.histogram(t, density=True) 9 plt.plot(hx[:-1].detach(), hy.detach()) 10 legends.append(f'layer {i} (Tanh)') 11plt.legend(legends) 12plt.title('activation distribution');
activation distribution
怎么读:六条线对应六层 tanh,每条线是该层所有输出值的分布。 健康的网络:每层都是钟形分布,集中在 0 附近,std ≈ 0.65~0.75,饱和率 < 5%。
判读口诀
  • 双峰在 ±1 → 神经元死亡(saturated)
  • 尖峰在 0 → 等价线性(没有非线性的好处)
  • 平缓钟形覆盖 (-0.8, 0.8) → 健康

诊断 2:梯度分布图(gradient distribution)

1plt.figure(figsize=(20, 4)) 2legends = [] 3for i, layer in enumerate(layers[:-1]): 4 if isinstance(layer, Tanh): 5 t = layer.out.grad # 该层 tanh 输出的梯度 6 print(f'layer {i} (Tanh): mean {t.mean():+f}, std {t.std():e}') 7 hy, hx = torch.histogram(t, density=True) 8 plt.plot(hx[:-1].detach(), hy.detach()) 9 legends.append(f'layer {i} (Tanh)') 10plt.legend(legends) 11plt.title('gradient distribution');
gradient distribution
怎么读:六条线对应六层 tanh 输出的梯度分布。 健康的网络:所有层的梯度量级相近(std 在同一数量级),表示梯度从 loss 一路顺利流回到每一层。
判读口诀
  • 浅层梯度比深层 100 倍以上 → 梯度消失(vanishing gradient)
  • 浅层梯度比深层 100 倍以上 → 梯度爆炸(exploding gradient)
  • 各层梯度量级相近 → 健康

诊断 3:权重梯度分布图(weights gradient distribution)

1plt.figure(figsize=(20, 4)) 2legends = [] 3for i, p in enumerate(parameters): 4 t = p.grad 5 if p.ndim == 2: # 只看 2D 权重矩阵,跳过 bias 6 print(f'weight {tuple(p.shape)} | mean {t.mean():+f} | ' 7 f'std {t.std():e} | grad:data ratio {t.std()/p.std():e}') 8 hy, hx = torch.histogram(t, density=True) 9 plt.plot(hx[:-1].detach(), hy.detach()) 10 legends.append(f'{i} {tuple(p.shape)}') 11plt.legend(legends) 12plt.title('weights gradient distribution');
weights gradient distribution
怎么读:每条线是某一层权重矩阵的梯度分布。图例里 grad:data ratiograd.std / param.std,量化每层"权重相对自己的被更新强度"。 健康的网络:所有 ratio 在同一数量级,意味着每层学习节奏一致。
判读口诀
  • ratio 跨层差 < 100 倍 = 各层学习节奏一致
  • 某层 ratio 特别大 = 那层学得太快,可能过拟合该层
  • 某层 ratio 特别小 = 那层几乎不学,参数浪费

诊断 4:更新比曲线(update ratio over time)★ 最关键的一张

1plt.figure(figsize=(20, 4)) 2legends = [] 3for i, p in enumerate(parameters): 4 if p.ndim == 2: 5 plt.plot([ud[j][i] for j in range(len(ud))]) 6 legends.append('param %d' % i) 7plt.plot([0, len(ud)], [-3, -3], 'k') # these ratios should be ~1e-3, indicate on plot 8plt.legend(legends);
update ratio over time
怎么读:横轴是训练步数,纵轴是 log10((lr × grad).std / param.std)。每条曲线对应一个权重矩阵,黑色水平线 = -3 是健康基准。 健康的网络:所有曲线密集聚拢在 -3 附近,表示每一步每一层都按"参数本身的 0.1%"在更新。
判读口诀
  • 所有层都贴近 -3 → 完美,lr 选得对
  • 整体高于 -2 → lr 太大,会震荡或发散
  • 整体低于 -4 → lr 太小,训练慢
  • 跨层差异 > 1 数量级 → 初始化有问题(不是 lr 问题)
这个指标的工程价值
  • 比 lr finder 更可靠(lr finder 只看一开始,这个看全过程)
  • 比 loss 曲线更早暴露问题(loss 不降的原因 = ratio 太小)
  • LLM 训练标配(wandb 自动追踪每一层的这个值)

一份"60 秒体检"清单

每次实现新模型,按以下顺序做 60 秒检查:

#检查看什么健康基准
1初始 loss第一步 loss 数值log(N)
2激活值直方图tanh 输出分布钟形,std 0.6~0.8,饱和率 < 5%
3梯度直方图各层激活的梯度量级相近,差距 < 100 倍
4update/param ratio训练前 100 步log10 在 -3 附近,跨层一致

通过 → 安心训练;不通过 → 改初始化 / 加 Norm / 调 lr。

六、本节引入的核心训练技巧

这一节首次引入了"诊断式训练"的思维方式 —— 后面所有大模型训练都靠这个:

1. 手算初始 loss

  • 任何 N 分类网络的合理初始 loss = log(N)
  • 训练第一步 loss 远高于这个 → 初始化有问题
  • 这是 60 秒就能做的"网络体检"

2. Kaiming 初始化

  • 公式:W = randn(fan_in, fan_out) * gain / sqrt(fan_in)
  • 原理:保持每层前向方差不变
  • gain:tanh 5/3,ReLU √2,sigmoid 1.0
  • 目标:让信号既不爆炸也不消失

3. BatchNorm 的"强制修复"

  • 不管初始化好不好,每层都标准化到均值 0、方差 1
  • 训练态用 batch 统计,推理态用 running 统计
  • 加 BN 后对应的 Linear 不要 bias
  • 让深层网络的训练稳定性提升一个数量级

4. 激活值直方图 + 饱和率

  • 健康基准:tanh 饱和率 < 5%(视频里截图 15-29% 是修复前)
  • 健康基准:每层 std 在 0.6 ~ 0.8 之间
  • 直方图应该是钟形,不是双峰

5. update/param ratio

  • 计算:(lr × grad).std() / param.std()
  • 健康值:log10(ratio) ≈ -3
  • 这是调学习率的真正科学方法,比 lr finder 更可靠
  • 太小(< -4)→ lr 太低;太大(> -2)→ lr 太高

6. 模块化的 PyTorch 风格

  • 每个层有 __init____call__parameters() 三件套
  • 写自己的 Linear / Tanh / BatchNorm1d 就是 nn.Module 的雏形
  • 训练循环里用 for layer in layers: x = layer(x) 优雅地组合

七、值得记住的几个"坑"

  1. 加 BN 的层不要 bias
    Linear(..., bias=False) 然后接 BatchNorm1d。BN 的 xmean 减法会抵消任何 bias,不去掉只是多一个无意义参数。

  2. BN 的训练态/推理态切换
    训练时必须 self.training = True,用 batch 统计;推理时必须 self.training = False,用 running 统计。忘记切换 → 推理结果飘忽不定。PyTorch 里用 model.train() / model.eval() 切换。

  3. running_mean 的更新要在 torch.no_grad()
    否则会被 autograd 追踪,浪费显存而且语义不对。running_mean 不是可学参数,是 buffer。

  4. batch_size = 1 时 BN 崩坏
    单样本算方差是 0(或 NaN),除以 0 → 数值崩溃。这也是 LayerNorm 取代 BN 的原因之一。

  5. 饱和率高的真凶往往是初始化,不是 lr
    看到 tanh 饱和先查初始化。lr 太大只会让 loss 震荡,不会让初始就饱和。

  6. 诊断要在训练早期做,不是等 loss 不降才看
    第一步 forward 后就该看一次激活值直方图。训练前 100 步是判断初始化好坏的黄金窗口。

八、与后面 Transformer / LLM 的关系

这一节讲的是 2015 年的 BatchNorm,但你要带着"未来视角"理解:

Part 3 内容现代 LLM 用什么为什么变了
BatchNormLayerNorm / RMSNormLLM 推理 batch 多变;LayerNorm 不依赖 batch
tanhGELU / SwiGLUtanh 易饱和;GELU 平滑且不饱和
Kaiming initscaled init(按深度缩小)Transformer 残差累加方差,需更小初始化
饱和率诊断同样有效诊断思维永不过时
update/param ratio同样有效跨架构通用

最重要的连续性

  • BatchNorm 可能退役,但"为什么需要 normalization"的洞察永远适用
  • 诊断工具(激活值、梯度、ratio)跨任何架构都管用
  • 学会"看图诊断"是从"调参玄学"到"诊断科学"的分水岭

九、Day 16 之前的建议练习

  • 不看代码,自己默写出 Kaiming 初始化公式,并说出 fan_in 是什么
  • 手画 tanh 函数图,标出饱和区和线性区,说明为什么饱和是问题
  • 用一句话说清 BatchNorm 的 4 步操作
  • 解释 BN 训练态和推理态为什么不同
  • 在 Part 2 的代码里加上 Kaiming 初始化,对比训练曲线
  • 实现并加上 BatchNorm,看激活值直方图怎么变健康
  • 写一个函数,输入模型 + batch,输出每层的 mean / std / saturated %
  • 核心检验:给你一张激活值直方图,你能不能 30 秒内判断网络是健康/饱和/死亡

十、给 1B 模型训练的预演

你将来训练 1B 模型时,这一节学的所有东西都会以"升级版"出现:

Part 3 学到的1B 训练对应的
手算初始 loss = log(27) ≈ 3.29手算 log(50257) ≈ 10.83(GPT-2 词表)
Kaiming 初始化scaled init: std = 0.02 / sqrt(2 × num_layers)
BatchNormRMSNorm(LLaMA 用的)
tanh 饱和率GELU/SwiGLU 输出统计
update/param ratio ≈ 1e-3同样的基准,跨规模通用
激活值直方图wandb / tensorboard 自动追踪
核心心得:Part 3 表面上讲的是 "BatchNorm 和 tanh",本质上讲的是 "如何科学地训练一个深度网络"。这套方法论永远不过时 —— 工具会变(BN → LayerNorm),但"诊断思维"是终生受用的。

十一、一句话总结

Part 3 把你从"会跑模型"升级到"会诊断模型"。学完后你应该养成的习惯是:每次新模型先看每层的 mean/std/saturated%,第一步 loss 是否等于 log(N),update/param ratio 是否在 1e-3 附近 —— 这三个体检指标,是你从此以后训练任何网络的入门动作。