Makemore Part 5 讲义:Building a WaveNet
〇、本节地图与两天切法
一句话概括本节:把 Part 3 的平铺 MLP 重构成模块化代码,再把它改造成 WaveNet 式的层级结构,顺路抓一个不报错的 bug。Part 4 是微观(每个算子的梯度),Part 5 是宏观(网络怎么搭)。
| 时间(官方章节) | 内容 | 本讲义 | 难度 |
|---|---|---|---|
| 0:00 – 6:56 | 引言 + 起步代码走读 | 第一节 | ★ |
| 6:56 – 9:16 | 修 loss 曲线图 | 第三节 | ★ |
| 9:16 – 17:11 | PyTorch 化:层、容器 | 第二节 | ★★ |
| 17:11 – 21:36 | WaveNet 概览 + block_size 3→8 | 第四节 | ★★ |
| 21:36 – 37:41 | 实现层级结构(本节核心) | 第四、五节 | ★★★ |
| 37:41 – 46:07 | 首训 + BN bug + 修复重训 | 第六节 | ★★★ |
| 46:07 – 51:34 | 扩容 + 实验方法 + dilated conv | 第七、八节 | ★★ |
| 51:34 – 56:00 | torch.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 的约定):
- 构造函数里创建并保存参数
__call__(x)做 forward,结果存self.out(方便事后检查每层输出——第六节抓 bug 全靠它)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)
现在动架构。目标结构——一棵二叉树:
每一层做的事完全相同:把相邻两个片段的表示拼在一起,过一个 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
六、事故现场: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——提升不大,但这不是重点。重点是:
① 打印形状——每层 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 收尾时把自己的工作流明说了,逐条展开:
- 盯住形状:动手写层之前先在纸上推完整形状链(练习 STEP 4 的表格就是在练这个);写完用形状追踪打印核对。大部分 bug 死在这一步之前。
- 先跑小的:结构性改动先在小配置上验证(能过拟合小数据、loss 曲线形态正常),确认语义无误再放大。永远不要在大配置上调试。
- 对照参考实现:自己写的层和 torch.nn 的同名层行为比对(输出、buffer 形状、eval 行为)。BN bug 就是这么抓到的。
- 一次只改一个变量:加长上下文、换层级结构、修 bug、扩容——每步单独重训记录 loss(第七节那张表就是这么来的)。多个改动混在一起,你永远不知道是哪个起了作用。这就是 46:58 说的"实验 harness"的雏形,也是以后做消融实验(ablation)的纪律。
十、预备问题与自测
看视频前带着这些问题
- block_size 从 3 到 8,平铺结构的第一层 Linear 参数量怎么变?
- 一个"层"的类需要哪三样东西?running_mean 为什么不放进 parameters()?
- (32, 4, 20) @ (20, 68) 的结果形状?矩阵乘发生在哪个维度、批维是谁?
- 为什么 view 能实现"拼接相邻位置"而不用任何计算?
- BN 收到 3D 输入时,"正确"的统计量形状应该是什么?
学完后的自我检测(答不上 3 道以上回去重看对应小节)
| # | 题目 | 要点(对应小节) |
|---|---|---|
| 1 | 平铺结构的两大天花板 | 一口闷的信息瓶颈;参数线性膨胀(一) |
| 2 | 默写 FlattenConsecutive 的 __call__,解释 view 为什么等于拼接 | 内存连续排布 + 重画格线(五) |
| 3 | 层级网络中第二个 Linear 的 fan_in 及理由 | 136 = 68×2,FC(2) 拼接了通道(五) |
| 4 | BN bug 的根因 / 症状 / 为什么不报错 / 修法 | 写死 dim=0;(1,4,68);broadcasting 宽容;dim=(0,1)(六) |
| 5 | BN 前 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) 张量过 Linear | Transformer 全程的数据形状约定 |
| BN 沉默 bug + 三条防御 | "先小规模验证语义,再烧卡"的铁律 |
| 单变量实验纪律 | 消融实验、超参扫描的方法论雏形 |
| 渐进融合、感受野逐层扩大 | conv → attention 的思想桥梁:Part 6 只是把"固定配对"换成"动态配对" |