模型与特征可视化
本教程共 60 篇 · 第 31 篇 · 更新于 2026-08-17 · 约 4 分钟阅读
本节目标:学会三个层面的可视化——看数据批次、看中间特征图、看梯度流向,并了解 CAM 这种「模型注意力」可视化。
给网络拍 X 光
上一章用 TensorBoard 看训练的外部表现。这一章换个视角,看看网络内部:数据进去后变成了什么样?梯度怎么流动?模型到底盯着图片的哪里看?
这类可视化像给网络拍 X 光。模型不是黑盒,你完全可以把中间结果拿出来看。手段就两个:拼图展示,和钩子(hook)截取。
看数据:make_grid 拼图
训练前先看看喂进去的数据长什么样。一批图片的形状是 (N, C, H, W),直接没法显示。torchvision.utils.make_grid 能把它们拼成一张大图:
import torch
import torchvision.utils as vutils
import matplotlib.pyplot as plt
batch = torch.rand(16, 3, 64, 64) # 16 张 64x64 的图
grid = vutils.make_grid(batch, nrow=4, padding=2)
plt.imshow(grid.permute(1, 2, 0))
plt.axis("off")
plt.show()
nrow=4 表示每行放 4 张。真实场景里把 batch 换成 DataLoader 取出的一批图就行。想在图上标点,比如人脸关键点,再叠加 plt.scatter 画散点即可。
看特征图:hook 截取中间输出
数据经过第一层卷积后,变成了一批特征图(feature map),每张图代表一种图案的响应。想看它们,就要在前向传播途中「截胡」。PyTorch 的钩子机制干的就是这件事:
import torch
import torch.nn as nn
import matplotlib.pyplot as plt
net = nn.Sequential(
nn.Conv2d(3, 16, 3, padding=1),
nn.ReLU(),
nn.Conv2d(16, 32, 3, padding=1),
nn.ReLU(),
)
feats = {}
def hook_fn(module, input, output):
feats["conv1_out"] = output.detach()
net[0].register_forward_hook(hook_fn)
x = torch.rand(1, 3, 64, 64)
net(x)
f = feats["conv1_out"] # 形状 [1, 16, 64, 64]
fig, axes = plt.subplots(4, 4, figsize=(6, 6))
for i, ax in enumerate(axes.flat):
ax.imshow(f[0, i].numpy(), cmap="gray")
ax.axis("off")
plt.show()
register_forward_hook 注册的函数会在该层前向结束后被调用,output 就是这层的输出。画出来的 16 张图里,有的响应边缘,有的响应纹理,各司其职。特征图是网络对输入的「解读」,多看几层,你对卷积的理解会扎实很多。
Tip特征图数值可能很小甚至为负,直接画会是一片黑。画之前先归一化:
img = (f[0, i] - f[0, i].min()) / (f[0, i].max() - f[0, i].min() + 1e-6)。
Note
detach()很关键。钩子里拿到的输出连着计算图,不 detach 直接画图,会拖着整张计算图占内存。
顺带检查形状
hook 的另一个日常用途:打印每层输入输出的形状,排查形状错误。不用改模型代码:
# 沿用上面的 net 和 x
def shape_hook(module, input, output):
print(module.__class__.__name__, "->", output.shape)
for layer in net:
layer.register_forward_hook(shape_hook)
net(x) # 每一层的输出形状一目了然
模型搭好后跑一次,形状跟预期对不对,立刻见分晓。数据维度传错的报错,九成能靠这个提前发现。
看梯度:信号流到哪去了
特征图看「前向」,梯度看「反向」。梯度太小,前面的层学不动;梯度太大,训练会炸。给中间输出挂个钩子,就能收集每层梯度的平均绝对值:
import torch
import torch.nn as nn
import matplotlib.pyplot as plt
grads = []
def collect_grad(grad):
grads.append(grad.abs().mean().item())
def fwd_hook(module, input, output):
output.register_hook(collect_grad)
net = nn.Sequential(
nn.Linear(784, 128), nn.Sigmoid(),
nn.Linear(128, 128), nn.Sigmoid(),
nn.Linear(128, 128), nn.Sigmoid(),
nn.Linear(128, 10),
)
handles = []
for layer in net:
handles.append(layer.register_forward_hook(fwd_hook))
x = torch.randn(8, 784)
loss = net(x).sum()
loss.backward()
plt.plot(grads[::-1], marker="o")
plt.xlabel("层(从浅到深)")
plt.ylabel("平均 |梯度|")
plt.show()
for h in handles:
h.remove()
梯度是反向传播时从后往前记录的,所以画图前要倒序,才能对应「浅层到深层」。这张折线图就是常说的「梯度流」图。
为什么要关心这个?训练的本质是让每层参数都收到合适的更新信号。浅层梯度消失,前面的层就永远学不动,整个网络退化成「只有后面几层在干活」。梯度爆炸则相反,一次更新就把参数甩到千里之外。所以画梯度流图,是排查「模型学不动」类问题时的第一站。
Note钩子里拿到的梯度可能是 None。张量不需要梯度(requires_grad=False),或者计算图被 detach 中断,都不会有梯度。画图前记得判空。
官方教程做过一个对比实验:同样层数的网络,加了批量归一化(BatchNorm)的,各层梯度都维持在非零水平;没加的,浅层梯度迅速趋近于零。这就是梯度消失的可视化证据。第 28 章讲过 BatchNorm,这里的实验正好是它的「疗效图」。
顺带看看激活
除了梯度,激活值的分布也能讲故事。ReLU 之后的特征图往往一大半是 0,说明不少神经元在「偷懒」。把某层输出的均值、非零比例打印出来,是判断网络是否健康的 quick check:
# 沿用上一节的 net、x 和 feats
act = feats["conv1_out"]
print("均值:", act.mean().item(), "非零比例:", (act != 0).float().mean().item())
非零比例过低,可能初始化把信号弄没了;过高且均值很大,警惕数值溢出。
权重本身也能看
参数张量也能直接画成图。卷积核尺寸小,适合一张张看:
# 沿用特征图一节的 net(第一个卷积层)
w = net[0].weight.detach() # 形状 [16, 3, 3, 3]
w = (w - w.min()) / (w.max() - w.min()) # 归一化到 [0, 1]
w = w.permute(0, 2, 3, 1) # 转成 HWC
plt.imshow(vutils.make_grid(w, nrow=4).permute(1, 2, 0))
plt.axis("off")
plt.show()
训练前看卷积核,是随机的噪声图案;训练后往往能看出边缘、色块等结构,说明网络学到了东西。这算是「一眼看出模型有没有学进去」的土办法。
可视化要注意的坑
画图路上有几个坑,先踩为敬:
- GPU 上的张量要先
.cpu()再.numpy(),直接转会报错。 - 通道顺序:PyTorch 是 C×H×W,matplotlib 要 H×W×C,记得
permute(1, 2, 0)。 - 做过 Normalize 的图像数值会偏到负区间,画之前先反归一化,否则一片灰。
Tip想省事的话,前三个坑 TensorBoard 的
add_image都替你处理了。本地快速看用 matplotlib,长期记录用 TensorBoard。
Warning钩子用完记得
handle.remove()移除。反复注册不清理,会越攒越多,拖慢训练还占内存。
CAM:看模型盯着哪里看
还有一个很有意思的可视化,叫类激活图(Class Activation Map,CAM)。它的升级版 Grad-CAM 思路很妙:把最后一层卷积的特征图,按各自对目标类别的梯度加权求和,得到一张热力图。热力图亮的地方,就是模型做判断时重点看的区域。
比如分类「猫」,热力图如果亮在猫脸上,说明模型真的在看猫;如果亮在背景上,说明它在投机取巧。这是排查模型错误的有力工具。
经典 CAM 有个限制:要求最后一层接全局平均池化,网络结构被绑死。Grad-CAM 改用梯度加权,任何 CNN 都能用,所以更流行。PyTorch 本身没有内置 Grad-CAM,需要借助第三方库实现。原理懂了,用起来就是套公式的事。
小结
这一章给网络拍了三张 X 光:拼图看数据,hook 看特征图和梯度流,CAM 看注意力区域。下一章把镜头对准性能——用 Profiler 给模型做体检,找出最慢的那一步。