Stephen 技术博客

AI和数据平台的工程实践

← 返回文章列表

Makemore Part 5 讲义:Building a WaveNet

对应 Karpathy "Building makemore Part 5: Building a WaveNet" · 时长约 56 分钟
视频链接:https://www.youtube.com/watch?v=t3YJ5hKiMQ0 · 学习日:Day 20–21
配套练习:Makemore_Part5_实战_WaveNet_练习.ipynb(+ 参考答案)
🌊 从"更深"到"有结构地深"

〇、本节地图与两天切法

一句话概括本节:把 Part 3 的平铺 MLP 重构成模块化代码,再把它改造成 WaveNet 式的层级结构,顺路抓一个不报错的 bug。Part 4 是微观(每个算子的梯度),Part 5 是宏观(网络怎么搭)。

时间(官方章节)内容本讲义难度
0:00 – 6:56引言 + 起步代码走读第一节
6:56 – 9:16修 loss 曲线图第三节
9:16 – 17:11PyTorch 化:层、容器第二节★★
17:11 – 21:36WaveNet 概览 + block_size 3→8第四节★★
21:36 – 37:41实现层级结构(本节核心)第四、五节★★★
37:41 – 46:07首训 + BN bug + 修复重训第六节★★★
46:07 – 51:34扩容 + 实验方法 + dilated conv第七、八节★★
51:34 – 56:00torch.nn、开发流程、展望第九节

Day 20:看 0:00–37:41,读本讲义第一~五节,做练习 STEP 1–4。Day 21:看 37:41–结尾,读第六~九节,做 STEP 5–6 + 自测。

一、出发点:平铺 MLP 撞到了什么天花板

Part 3/4 的网络长这样:3 个字符各查 10 维 embedding,拼成 30 维,一层隐藏层,输出 27 类。200k 步训练后 val loss ≈ 2.10。想更好,最自然的想法是看更长的上下文——3 个字符实在太少了,预测 "…avi?" 的下一个字母时,模型连名字开头是什么都不知道。

把 block_size 从 3 改到 8,平铺结构立刻暴露两个问题:

问题 1:第一层被迫"一口闷"

8 个字符的 embedding 拼成 80 维,一层就要把 8 个字符的全部信息压进 hidden 向量。信息在网络的第一步就被挤过最窄的瓶颈——前面的字符和后面的字符、相邻的和相隔的,全部一视同仁地搅在一起。网络没有机会先弄清"局部"再理解"整体"。

问题 2:参数花在错误的地方

第一层 Linear 的参数量 = 80 × n_hidden,随上下文长度线性膨胀,而且这些参数都花在"一步到位的大杂烩"上。上下文再翻倍(16、32 个字符)这条路就走死了。

本节的答案(借自 DeepMind 2016 年的 WaveNet 论文,原本用于语音生成):不要一口闷,两两融合、逐层聚合——8 个字符先两两合成 4 个"双字符组",再合成 2 个"四字符组",最后合成 1 个"八字符表示"。Karpathy 的说法:squash the information slowly(慢慢挤压信息)。

但在动架构之前,得先解决一个工程问题:Part 3 的代码全是散装张量(裸的 W1、b1、手写的 forward 行),每改一次结构都要大动干戈。所以本节的前半场是——

二、第一件事:把代码 PyTorch 化(9:16)

目标:把每种运算封装成,遵守三条 API 约定(正是 torch.nn 的约定):

  1. 构造函数里创建并保存参数
  2. __call__(x) 做 forward,结果存 self.out(方便事后检查每层输出——第六节抓 bug 全靠它)
  3. parameters() 返回该层的可训练参数列表

2.1 Linear —— 带初始化纪律的矩阵乘

class Linear:
    def __init__(self, fan_in, fan_out, bias=True):
        self.weight = torch.randn((fan_in, fan_out), generator=g) / fan_in**0.5
        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])

两个细节都是前几节的直接回收:/ fan_in**0.5 是 Part 3 的方差保持初始化(不做它,深网络的激活值逐层放大或缩小);bias=False 选项是给 BN 前的 Linear 用的——Part 4 里你亲手算出 db1 ≈ 0:BN 第一步减均值,bias 加多少就被减掉多少,纯属浪费参数。GPT、LLaMA 里所有 Norm 前的 Linear 都是 bias=False,源头就是这个。

2.2 BatchNorm1d —— 第一次有了"训练/推理"两副面孔

class BatchNorm1d:
    def __init__(self, dim, eps=1e-5, momentum=0.1):
        self.eps, self.momentum = eps, momentum
        self.training = True
        self.gamma = torch.ones(dim)     # 可训练:缩放
        self.beta  = torch.zeros(dim)    # 可训练:平移
        self.running_mean = torch.zeros(dim)   # 不可训练的 buffer
        self.running_var  = torch.ones(dim)

    def __call__(self, x):
        if self.training:
            xmean = x.mean(0, keepdim=True)   # ⚠️ 这行埋着本节的大 bug,第六节揭晓
            xvar  = x.var(0, keepdim=True)
            with torch.no_grad():             # buffer 更新不参与求导
                self.running_mean = (1-self.momentum)*self.running_mean + self.momentum*xmean
                self.running_var  = (1-self.momentum)*self.running_var  + self.momentum*xvar
        else:
            xmean, xvar = self.running_mean, self.running_var
        self.out = self.gamma * (x - xmean) / torch.sqrt(xvar + self.eps) + self.beta
        return self.out

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

比 Part 3 多出来的东西全是"工程":training 开关(训练时用当前 batch 的统计量,推理时一个样本没有"batch"可言,用训练期攒下的 running 统计量);momentum 动量更新(每步把 running 值向当前 batch 值挪 10%,等效于指数滑动平均,训练结束时它≈全数据集的统计量,省掉 Part 4 Ex4 里那次"全量校准");参数 vs buffer 的区分(gamma/beta 有梯度要训练,running_mean/var 没有梯度只是记账——PyTorch 里对应 register_buffer)。

2.3 其余三个小件 + 容器

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

class Embedding:                      # Part 2 的 C[Xb] 穿上类的外衣
    def __init__(self, num_embeddings, embedding_dim):
        self.weight = torch.randn((num_embeddings, embedding_dim), generator=g)
    def __call__(self, IX):
        self.out = self.weight[IX]    # 整数索引查表:(B,T) → (B,T,C)
        return self.out
    def parameters(self): return [self.weight]

class Sequential:                     # 容器:把层串成流水线
    def __init__(self, layers): self.layers = layers
    def __call__(self, x):
        for layer in self.layers:
            x = layer(x)
        self.out = x
        return self.out
    def parameters(self):
        return [p for layer in self.layers for p in layer.parameters()]

重构的回报立竿见影——训练循环从此只依赖两个接口,换任何架构都不用改训练代码

logits = model(Xb)                    # 接口 1:model(x)
loss = F.cross_entropy(logits, Yb)
...
for p in model.parameters():          # 接口 2:parameters()
    p.data += -lr * p.grad

三、插曲:修好 loss 曲线(6:56)

此前各节的 loss 曲线是一团毛刺,原因不是训练有问题,而是每一步画的是单个 minibatch 的 loss——batch 只有 32 个样本,抽到"简单批"就低、"难批"就高,噪声淹没了趋势。修法一行:

plt.plot(torch.tensor(lossi).view(-1, 1000).mean(1))   # 每 1000 步取平均再画

view(-1, 1000) 把一维 loss 序列折成"每行 1000 步"的二维表,mean(1) 按行取均值。曲线立刻变得光滑可读——以后你训 1B 模型看的 loss 曲线(wandb 里的 smoothing)就是这个思想。

四、核心:WaveNet 的渐进融合(17:11 – 37:41)

现在动架构。目标结构——一棵二叉树:

c1 c2 c3 c4 c5 c6 c7 c8 c1c2 c3c4 c5c6 c7c8 c1..c4 c5..c8 c1..c8 → 预测

每一层做的事完全相同:把相邻两个片段的表示拼在一起,过一个 Linear + BN + Tanh,产出一个"更大片段"的表示。第一层学二元组("ar"、"nn" 这类字符搭配),第二层学四元组(音节级),第三层学八元组(整个名字片段)。低级模式先各自成形,再组合成高级模式——而不是像平铺那样第一层就把所有东西搅成一锅。

关键理解:批维度的"变形记"

实现的巧妙之处在于:不需要写任何循环,只需要玩转形状。回想 Part 3,(32, 30) 的张量过 Linear 时,32 是"陪跑"的批维度。现在把张量变成三维 (32, 4, 40)——PyTorch 的矩阵乘会把前面所有维度都当批维,只对最后一维做线性变换:

(32, 4, 40) @ (40, 68) → (32, 4, 68)     # 等于 32×4=128 个 40 维向量并行过同一个 Linear

也就是说,"4 个双字符组各自过一遍 Linear"这件事,靠把组数塞进批维度就自动并行了,且四个组共享同一套权重(同一个 Linear)。这就是权重共享——卷积的本质,第八节还会回来说它。

五、灵魂零件:FlattenConsecutive 逐行剖析

整个改造只需要一个新零件——负责"把相邻 n 个位置拼在一起"的层:

class FlattenConsecutive:
    def __init__(self, n):
        self.n = n
    def __call__(self, x):
        B, T, C = x.shape                    # 批,时间(位置数),通道
        x = x.view(B, T//self.n, C*self.n)   # ① 相邻 n 个位置的通道并排拼接
        if x.shape[1] == 1:
            x = x.squeeze(1)                 # ② 时间维只剩 1 时挤掉,回到 2D
        self.out = x
        return self.out
    def parameters(self): return []

①为什么一个 view 就能"拼接相邻位置"?因为张量在内存里是按 (B, T, C) 的顺序连续排布的:位置 t=0 的 10 个通道后面紧跟着 t=1 的 10 个通道。view(B, T//2, C*2) 只是重新画格线:原来"每 10 个数一格",现在"每 20 个数一格"——恰好把相邻两个字符的向量框进同一格。零拷贝、零计算(Part 4 学过:view 的 backward 也只是 reshape 回去)。

②的 squeeze 处理最后一层:融合到只剩 1 个"片段"时,(32, 1, 136) 挤成 (32, 136),好接最后的普通 Linear。(顺带:Part 3 的一次性压扁 = FlattenConsecutive(block_size) 的特例。)

完整形状链(务必亲手推一遍——练习 STEP 4 的形状预测就是它)

输入 Xb                          (32, 8)          8 个字符的整数索引
Embedding(27, 10)                (32, 8, 10)
FlattenConsecutive(2)            (32, 4, 20)      8 字符 → 4 个双字符组
Linear(20→68) + BN + Tanh        (32, 4, 68)      4 组并行过同一套权重
FlattenConsecutive(2)            (32, 2, 136)     4 组 → 2 个四字符组
Linear(136→68) + BN + Tanh       (32, 2, 68)
FlattenConsecutive(2)            (32, 136)        2 组 → 1 个八字符表示(squeeze!)
Linear(136→68) + BN + Tanh       (32, 68)
Linear(68→27)                    (32, 27)         logits
最容易写错的地方:FlattenConsecutive 之后的 Linear,fan_in 是上一层通道数 × 2(20、136、136),不是 68。写错了 PyTorch 会在矩阵乘处报形状错——这是"好 bug",会喊疼。真正危险的是下一节这种不喊疼的。

六、事故现场:BatchNorm 的沉默 bug(38:50)

网络组装好,训练跑通,loss 正常下降——一切看起来都好。然后 Karpathy 做了一个随手的检查(这个习惯本身就是本节的教学内容):打印每层的输出形状和 BN 的统计量形状。

>>> model.layers[3].running_mean.shape
torch.Size([1, 4, 68])          # ???

破案过程

它应该是什么:BN 的本意是给 68 个通道维护一套均值/方差,即形状 (68,) 或 (1, 1, 68)。

为什么变成了 (1, 4, 68):我们的 BN 写死了 x.mean(0, keepdim=True)。2D 输入 (32, 68) 时没问题——对 dim 0(批维)归约,得 (1, 68)。但层级网络的中间 BN 收到的是 3D 输入 (32, 4, 68):只对 dim 0 归约,剩下 (1, 4, 68)——4 个时间位置各自算了一套统计量

错在哪:语义上,4 个位置的向量过的是同一个 Linear(权重共享),它们的输出属于同一个分布,本该把 32×4=128 个向量放在一起统计。分开统计的后果:每套统计量只基于 32 个样本(噪声更大),推理时同一个通道在不同位置还会被不同的均值/方差归一化(语义错乱)。

为什么不报错:因为后续的 (x - xmean) / … 全靠 broadcasting——(32,4,68) 减 (1,4,68) 完全合法。广播机制太"宽容"了,形状能对上就往下算。损失照降、采样照出,只是一切都悄悄差了一点。

修复:一行

dim = 0 if x.ndim == 2 else (0, 1)      # 3D 输入:批维和时间维一起归约
xmean = x.mean(dim, keepdim=True)        # → (1, 1, 68) ✓
xvar  = x.var(dim, keepdim=True)

修复后重训,val loss 从 ≈2.029 到 ≈2.022——提升不大,但这不是重点。重点是:

这一段是全视频最值钱的 8 分钟。深度学习的 bug 大多长这样:不抛异常、指标还行、模型"基本能用"。防御手段只有三个,全是习惯而非技巧:
打印形状——每层 self.out 的形状、每个 buffer 的形状,跟你纸上推的对照(练习里的形状追踪 cell);
小数据先过拟合——几十个样本 loss 能压到接近 0,说明语义大体是对的;
对照参考实现——Karpathy 就是对照 torch.nn.BatchNorm1d 的文档发现自己的行为和官方不一致的。
训练 1B 模型时,这样一个 bug 可能烧掉几千美元 GPU 时间才被察觉。

七、扩容与战绩(46:07)

模块化的第二个回报:扩容只改两个数字。n_embd 10→24, n_hidden 68→128,参数量 22,397 → 76,579,训练循环一个字不动。

模型参数量val loss(200k 步,视频口径)
Part 3/4 平铺 MLP(block=3)~22k≈ 2.10
平铺 MLP,block=8~22k≈ 2.02 ← 光加长上下文就有收益
WaveNet 层级(BN bug 在身)22,397≈ 2.029
WaveNet 层级(bug 修复)22,397≈ 2.022
WaveNet 扩容76,579≈ 1.993 🎉

值得注意的诚实结论:同参数量下层级结构相对平铺的提升并不惊艳(2.02 vs 2.029/2.022)。它的真正价值在于可扩展性——上下文再翻倍时,平铺的第一层参数线性爆炸,层级结构只多加一轮融合(log 级增长)。架构的优劣要在规模变大时才显形,这也是大模型时代反复上演的剧本。

八、彩蛋:这和真正的卷积什么关系(47:44)

一句话定义卷积:同一组权重,沿着输入滑动,每个窗口算一次(滑动窗口 + 权重共享)。现在回看我们的树状结构:每层用同一个 Linear 处理所有相邻对——这就是 kernel size = 2 的一维卷积;三层的"跳距"分别是 1、2、4——这就是 WaveNet 论文里的 dilated causal convolution(空洞因果卷积):

  • causal(因果):每个输出只依赖它之前的输入(我们预测下一个字符,天然因果)
  • dilated(空洞):逐层加大间隔(1, 2, 4, 8…),感受野指数增长,log 层数覆盖任意长上下文

所以我们已经搭出了卷积网络的计算图,只是用 view + 批维技巧实现,没用 nn.Conv1d。真正的 Conv 层额外提供的只是效率:滑动窗口间的重叠计算可以复用(想象对一整句话的每个位置都做预测时,树的下层节点被多个上层共享)。卷积 = for 循环的高效实现,不是新的数学。

放进大图景:Part 5 的层级融合是"固定配对"——c3 永远只和 c4 融合。Part 6 的 attention 把它升级为"动态配对"——每个位置自己决定看谁、看多少。这一步之隔,就是 ConvNet 时代和 Transformer 时代的分界线。

九、方法论:深网络开发四步(52:28)

Karpathy 收尾时把自己的工作流明说了,逐条展开:

  1. 盯住形状:动手写层之前先在纸上推完整形状链(练习 STEP 4 的表格就是在练这个);写完用形状追踪打印核对。大部分 bug 死在这一步之前。
  2. 先跑小的:结构性改动先在小配置上验证(能过拟合小数据、loss 曲线形态正常),确认语义无误再放大。永远不要在大配置上调试。
  3. 对照参考实现:自己写的层和 torch.nn 的同名层行为比对(输出、buffer 形状、eval 行为)。BN bug 就是这么抓到的。
  4. 一次只改一个变量:加长上下文、换层级结构、修 bug、扩容——每步单独重训记录 loss(第七节那张表就是这么来的)。多个改动混在一起,你永远不知道是哪个起了作用。这就是 46:58 说的"实验 harness"的雏形,也是以后做消融实验(ablation)的纪律。

十、预备问题与自测

看视频前带着这些问题

  1. block_size 从 3 到 8,平铺结构的第一层 Linear 参数量怎么变?
  2. 一个"层"的类需要哪三样东西?running_mean 为什么不放进 parameters()?
  3. (32, 4, 20) @ (20, 68) 的结果形状?矩阵乘发生在哪个维度、批维是谁?
  4. 为什么 view 能实现"拼接相邻位置"而不用任何计算?
  5. BN 收到 3D 输入时,"正确"的统计量形状应该是什么?

学完后的自我检测(答不上 3 道以上回去重看对应小节)

#题目要点(对应小节)
1平铺结构的两大天花板一口闷的信息瓶颈;参数线性膨胀(一)
2默写 FlattenConsecutive 的 __call__,解释 view 为什么等于拼接内存连续排布 + 重画格线(五)
3层级网络中第二个 Linear 的 fan_in 及理由136 = 68×2,FC(2) 拼接了通道(五)
4BN bug 的根因 / 症状 / 为什么不报错 / 修法写死 dim=0;(1,4,68);broadcasting 宽容;dim=(0,1)(六)
5BN 前 Linear 为什么 bias=False减均值抵消 bias,Part 4 的 db1≈0(二)
6"卷积是什么"的一句话版本 + 我们的树对应什么滑动窗口+权重共享;dilation 1/2/4 的因果卷积(八)
7同参数量下层级结构提升不大,它的价值在哪可扩展性:上下文翻倍时参数 log 级 vs 线性(七)
8复述开发四步,并说出每步在本节的实例形状/小模型/对照/单变量(九)

十一、连回 1B 主线

Part 5 学的1B 训练里的
手写 nn.Module 体系(参数/buffer/training 开关)读 GPT/LLaMA 源码零障碍;理解 model.eval() 到底切换了什么
批维技巧:(B, T, C) 张量过 LinearTransformer 全程的数据形状约定
BN 沉默 bug + 三条防御"先小规模验证语义,再烧卡"的铁律
单变量实验纪律消融实验、超参扫描的方法论雏形
渐进融合、感受野逐层扩大conv → attention 的思想桥梁:Part 6 只是把"固定配对"换成"动态配对"