Appearance
实战第 2 步:模型设计(并故意踩第 1 个坑)
本篇定义网络结构。我们会先"按新手最自然的写法"写出一版有问题的代码,亲自体验一次 PyTorch 最高频的报错,然后学会从报错信息里读出修复方案——这比直接给你正确代码有价值得多。
先创建 model.py
python
# model.py
import torch.nn as nn
class DigitMLP(nn.Module):
"""3 层全连接网络:784 → 128 → 64 → 10"""
def __init__(self):
super().__init__()
self.fc1 = nn.Linear(28 * 28, 128)
self.relu = nn.ReLU()
self.fc2 = nn.Linear(128, 64)
self.fc3 = nn.Linear(64, 10) # 输出 10 个 logits,不加 Softmax
def forward(self, x):
return self.fc3(self.relu(self.fc2(self.relu(self.fc1(x)))))看起来没毛病?把它接到 train.py 里试一下——注意,坑来了。
踩坑时刻:形状不匹配
在 train.py 追加:
python
from model import DigitMLP
if __name__ == "__main__":
model = DigitMLP()
images, labels = next(iter(train_loader))
print("进入模型前:", tuple(images.shape))
out = model(images) # ← 就这一行,炸了
print("模型输出:", tuple(out.shape))运行 python train.py(终端),得到:
text
进入模型前: (64, 1, 28, 28)
RuntimeError: mat1 and mat2 shapes cannot be multiplied (1792x28 and 784x128)排查思路(重点学这个,不是背答案)
- 读报错:
mat1 and mat2 shapes cannot be multiplied (1792x28 and 784x128)。mat1是它眼中的输入、mat2是权重转置。它在说:想把[1792, 28]乘上[784, 128],中间维度 28 ≠ 784,乘不了。 - 等等,输入不是
[64,1,28,28]吗,1792 和 28 哪来的?nn.Linear只认"最后一维是特征":它把[64,1,28,28]悄悄压成"64×1×28 = 1792 个、特征维 28 的样本"。而fc1的权重是[128, 784],转置后[784, 128]——它期望每个样本特征维是 784,实际收到的是 28。根源永远是:喂给 Linear 的最后一维,和它的 in_features 不等。 - 对照设计意图:我们要的是每张图展平成 784 维,即
[64, 784]。中间两个维度(1 和 28)必须被压掉,而当前模型里没有任何一层做这件事。
一句话诊断:图片没展平就进了全连接层。
修复:显式加 Flatten 层
改 model.py,这是最清晰的方案:
python
# model.py —— 修复版
import torch.nn as nn
class DigitMLP(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(
nn.Flatten(), # [64,1,28,28] → [64,784],batch 维不动
nn.Linear(28 * 28, 128),
nn.ReLU(),
nn.Dropout(0.2),
nn.Linear(128, 64),
nn.ReLU(),
nn.Linear(64, 10),
)
def forward(self, x):
return self.net(x)再次运行 python train.py:
text
进入模型前: (64, 1, 28, 28)
模型输出: (64, 10)输出 [64, 10],与上一篇文章立下的"规格书"完全一致,模型层打通。
这次踩坑带走的 3 条经验
- 从
nn.Linear的 in_features 出发倒推形状:数据到它手上之前,最后一维必须恰好等于 in_features。在层与层之间随手print(x.shape)是最有效的排错手段。 - 报错信息里的两个形状都要读:
(1792x28 and 784x128)直接告诉了你"它以为的输入"和"它期望的权重",90% 的形状问题看一眼就能定位。 - 输出层不加 Softmax:
CrossEntropyLoss内部会做 LogSoftmax,输入必须是原始 logits。自己先套一层 Softmax 等于做了两次 softmax,数值语义全错,模型怎么也训不好。
模型通了,接下来写训练循环 → 第 3 步:训练循环(那里还有第 2 个坑等着你)。