
PyTorch 实战可视化神经网络中间层输出这次我们来看一个 PyTorch 实战中很常见、但网上零散资料比较多的需求可视化神经网络中间层输出。很多同学训练完模型后只知道准确率和 loss却说不清模型内部到底发生了什么。比如 ResNet 的 layer2 和 layer4 分别学到了什么特征为什么某个分类会出错某个卷积核是否已经“死掉”了这些问题都需要把中间层张量直接“捞出来”看。本文会先把实现原理讲清楚再给出一套可以直接跑的完整代码覆盖特征图提取、多通道绘制、输出分布统计和常见坑排查。全程以 CPU 跑通为目标有 GPU 会更快不需要高端显卡。1. 核心能力速览能力项说明任务类型深度学习模型调试 / 特征可视化 / 可解释性分析核心方法PyTorch 的register_forward_hook钩子机制依赖库Python、PyTorch、torchvision、matplotlib运行平台CPU 可以运行GPU 可加速输入格式任意尺寸图像代码会自动预处理输出内容指定层输出的特征图网格、通道热力图、张量分布直方图模型支持ResNet、VGG、CNN 自定义模型均可是否支持批量支持但批量大会放大内存占用适合场景模型调试、特征分析、论文配图、教学演示这套方案不依赖 TensorBoard也不依赖第三方可视化框架核心就是 PyTorch 自带的 hook 接口加 matplotlib 绘图灵活性很高。2. 可视化中间层输出的应用场景与使用边界2.1 适合解决什么问题模型调试训练 loss 不降时可以检查某一层输出是否退化成常量或全零判断是否出现梯度消失或“死亡 ReLU”。特征分析观察浅层网络关注纹理、边缘深层网络关注语义区域验证网络是否学到了有效特征。分类依据确认当模型给出错误分类时看最后一层卷积的特征图能直观看出模型到底“看”了图像的哪部分。论文和教学配图特征图可视化是深度学习课程、论文中最常见的示意图之一本文代码可以直接生成这种图。2.2 使用边界可视化中间层输出属于模型分析的常规手段但如果要公开模型的输入图片和特征图需要确认输入图片是否有合法授权涉及人脸、车辆、隐私场景的数据要打码或换用公开数据集。复现开源模型结果时注意模型权重和原项目 license。中间层特征图本身是模型内部张量不包含原始图像信息但高分辨率特征图重绘后仍可能与原图轮廓相似发布时要注意输入素材的版权。3. 本地部署环境准备可视化中间层输出本质上是一个 PyTorch 推理任务环境准备比训练简单很多。3.1 环境清单依赖项说明Python3.8 以上推荐 3.10PyTorch2.x 或 1.13 均可torchvision用于加载预训练模型和图像预处理matplotlib用于绘制特征图网格numpy张量转数组和统计计算3.2 安装命令使用 pip 安装核心依赖pip install torch torchvision matplotlib numpy如果使用 conda可以用conda create -n feat_vis python3.10 conda activate feat_vis pip install torch torchvision matplotlib numpy安装完成后验证环境是否正常python -c import torch; print(torch.__version__)3.3 硬件要求纯 CPU 即可完成本文全部操作只是推理速度稍慢。有 NVIDIA GPU 且安装了 CUDA 版 PyTorch 时前向传播会自动使用 GPU。不需要高显存显卡。只要不设置超大 batch普通 4G 显存也足够。4. 核心原理PyTorch Hook 机制4.1 为什么不能直接打印中间变量直接改模型 forward 函数里的print(x.shape)确实可以但工程上很麻烦你需要在每个想观察的层后面加打印代码改完模型结构还要改回来。PyTorch 提供了register_forward_hook可以在不修改模型源码的情况下在某个模块执行完 forward 之后自动回调一个函数。4.2 hook 的工作流程使用register_forward_hook的流程如下定义要观察的网络层模块。调用module.register_forward_hook(hook_fn)。hook_fn(module, input, output)会在该模块前向传播完成后自动执行。在hook_fn里把output保存到外部列表用于后续可视化。4.3 hook 函数的基本写法import torch import torch.nn as nn features {} def make_hook(name): def hook_fn(module, input, output): # output 是当前层前向传播的输出张量 # 对 BatchNorm、Dropout 等层output 类型可能不同一般只处理 Tensor if isinstance(output, torch.Tensor): features[name] output.detach().cpu() print(f[{name}] output shape: {list(output.shape)}) else: print(f[{name}] output type: {type(output)}) return hook_fn调用方式model torchvision.models.resnet18(pretrainedTrue) model.layer2[-1].register_forward_hook(make_hook(layer2_last))模型推理结束后features[layer2_last]就是一个形状为[B, C, H, W]的浮点张量。5. 完整代码提取并可视化 ResNet18 中间层特征图下面给出一套完整可运行的示例代码。使用 ResNet18 作为演示模型因为它的层结构清晰适合展示不同深度的特征差异。5.1 完整脚本import torch import torchvision import torchvision.transforms as transforms import numpy as np import matplotlib.pyplot as plt from PIL import Image # 1. 加载预训练模型 model torchvision.models.resnet18(weightstorchvision.models.ResNet18_Weights.DEFAULT) model.eval() # 2. 定义特征保存字典 feature_maps {} # 3. 注册 hook def register_hook(name): def hook_fn(module, input, output): if isinstance(output, torch.Tensor): feature_maps[name] output.detach() else: print(flayer {name} output is not a Tensor: {type(output)}) return hook_fn # 观察四个不同深度的层 model.conv1.register_forward_hook(register_hook(conv1)) model.layer1[-1].register_forward_hook(register_hook(layer1_last)) model.layer2[-1].register_forward_hook(register_hook(layer2_last)) model.layer4[-1].register_forward_hook(register_hook(layer4_last)) # 4. 图像预处理 transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) image Image.open(test.jpg).convert(RGB) input_tensor transform(image).unsqueeze(0) # [1, 3, 224, 224] # 5. 前向传播 with torch.no_grad(): output model(input_tensor) pred_class output.argmax(dim1).item() print(predicted class index:, pred_class) # 6. 可视化特征图 def show_feature_maps(tensor, num_cols8, title): 将 [C, H, W] 的特征图按网格绘制 C, H, W tensor.shape num_rows int(np.ceil(C / num_cols)) fig, axes plt.subplots(num_rows, num_cols, figsize(num_cols * 2, num_rows * 2)) axes axes.flatten() for i in range(num_cols * num_rows): if i C: feat tensor[i] # 归一化到 [0, 1] 便于显示 feat_min feat.min() feat_max feat.max() if feat_max - feat_min 1e-8: feat (feat - feat_min) / (feat_max - feat_min) axes[i].imshow(feat.numpy(), cmapviridis) axes[i].axis(off) else: axes[i].axis(off) fig.suptitle(title) plt.tight_layout() plt.show() # 7. 展示不同层的特征图 for name, tensor in feature_maps.items(): # 去掉 batch 维度得到 [C, H, W] feat_map tensor[0] print(f{name}: {list(feat_map.shape)}) # 显示第一个卷积层的特征图 show_feature_maps(feature_maps[conv1][0], num_cols8, titleconv1 output) # 显示 layer2 最后输出的前 16 个通道 show_feature_maps(feature_maps[layer2_last][0][:16], num_cols4, titlelayer2_last first 16 channels) # 显示 layer4 最后输出的前 16 个通道 show_feature_maps(feature_maps[layer4_last][0][:16], num_cols4, titlelayer4_last first 16 channels)5.2 运行结果预期运行后控制台会依次打印每一层的输出形状[conv1] output shape: [1, 64, 112, 112] [layer1_last] output shape: [1, 64, 56, 56] [layer2_last] output shape: [1, 128, 28, 28] [layer4_last] output shape: [1, 512, 7, 7]可以看到随着网络加深特征图数量从 64 增加到 512。空间分辨率从 112 缩小到 7。浅层特征图保留较多细节信息。深层特征图空间分辨率低更关注语义信息。6. 更实用的可视化技巧多通道特征图增强展示上面的代码有个问题当通道数很多时直接按顺序看前 16 个通道不一定能找到最有信息量的特征图。更实用的做法是按激活强度排序或绘制某张特征图的热力图叠加在原图上。6.1 筛选激活最强的通道对一张输入图片每个通道的特征图响应强度不同。可以计算每个通道的 L2 范数或均值选出响应最强的 K 个通道来展示。def select_top_k_channels(tensor, k16): 根据通道均值选出前 k 个通道 C, H, W tensor.shape scores tensor.view(C, -1).mean(dim1) # 每个通道的均值 top_k_indices scores.topk(min(k, C)).indices return tensor[top_k_indices], top_k_indices selected, indices select_top_k_channels(feature_maps[layer4_last][0], k16) print(selected channel indices:, indices.numpy()) show_feature_maps(selected, num_cols4, titlelayer4_last top-16 channels by mean activation)6.2 单通道热力图与原图叠加有时想看某个通道在图像上的响应区域可以把特征图放大到原图尺寸后叠加显示。from torch.nn.functional import interpolate def show_heatmap_overlay(original_img, feature_map, alpha0.5): original_img: PIL Image feature_map: [H, W] 单通道特征图 feat_resized interpolate( feature_map.unsqueeze(0).unsqueeze(0), size(original_img.height, original_img.width), modebilinear, align_cornersFalse ).squeeze() feat_resized feat_resized.numpy() feat_resized (feat_resized - feat_resized.min()) / (feat_resized.max() - feat_resized.min() 1e-8) plt.figure(figsize(8, 4)) plt.subplot(1, 2, 1) plt.imshow(original_img) plt.axis(off) plt.subplot(1, 2, 2) plt.imshow(original_img) plt.imshow(feat_resized, cmapjet, alphaalpha) plt.axis(off) plt.show() # 取 layer4_last 的第 0 个通道 show_heatmap_overlay(image, feature_maps[layer4_last][0][0])这种热力图叠加能清楚看到深层网络对图像哪个区域响应最强。6.3 卷积核可视化如果想看模型学到了什么卷积核可以绘制权重而不是输出特征图def visualize_kernels(kernel_tensor, num_cols8): kernel_tensor: [out_channels, in_channels, k, k] out_c kernel_tensor.shape[0] fig, axes plt.subplots( out_c // num_cols (1 if out_c % num_cols else 0), num_cols, figsize(num_cols * 1.5, (out_c // num_cols 1) * 1.5) ) axes axes.flatten() for i in range(out_c): # 对一个输出通道对输入通道求平均获得二维形状 kernel kernel_tensor[i].mean(dim0) kernel_min kernel.min() kernel_max kernel.max() kernel_norm (kernel - kernel_min) / (kernel_max - kernel_min 1e-8) axes[i].imshow(kernel_norm.numpy(), cmapgray) axes[i].axis(off) plt.suptitle(conv1 kernels) plt.show() visualize_kernels(model.conv1.weight.detach().cpu())7. 中间层输出分布统计除了直接看图量化统计中间层输出分布也很有价值。比如某层输出如果长期为 0说明 ReLU 可能已经“死亡”。7.1 输出均值和稀疏度def analyze_distribution(tensor, layer_name): feat tensor[0] # [C, H, W] values feat.numpy() zero_ratio (values 0).mean() print(f[{layer_name}]) print(f shape: {list(feat.shape)}) print(f mean: {values.mean():.6f}) print(f std: {values.std():.6f}) print(f min: {values.min():.6f}) print(f max: {values.max():.6f}) print(f zero ratio: {zero_ratio:.4f}) for name, tensor in feature_maps.items(): analyze_distribution(tensor, name)输出示例[conv1] shape: [64, 112, 112] mean: 0.018342 std: 0.541203 min: -2.213400 max: 4.071200 zero ratio: 0.3762如果某个深层特征图的 zero ratio 接近 1说明该层输出大量为 0模型表达能力可能已经不健康。7.2 绘制分布直方图def plot_histogram(tensor, layer_name, bins50): feat tensor[0].numpy().flatten() plt.figure(figsize(6, 4)) plt.hist(feat, binsbins, colorsteelblue, alpha0.8) plt.title(f{layer_name} output distribution) plt.xlabel(value) plt.ylabel(count) plt.grid(alpha0.3) plt.show() plot_histogram(feature_maps[layer2_last], layer2_last)8. 资源占用与性能观察8.1 hook 对推理速度的影响注册 hook 后每次前向传播会多执行一个 Python 回调函数。如果只在少数层注册速度影响可以忽略但如果对 ResNet 每一层都注册 hookPython 层回调调度会明显拖慢推理。建议只对需要观察的层注册 hook。分析完及时调用handle.remove()移除 hook。示例handle model.layer4[-1].register_forward_hook(hook_fn) # 用完移除 handle.remove()8.2 内存占用把所有中间层输出都保存在内存里内存占用会快速膨胀。例如输入[1, 3, 224, 224]ResNet18 的 layer3 输出是[1, 256, 14, 14]单张特征图占用约256 * 14 * 14 * 4 / 1024 / 1024 0.2MB。看起来不大但如果保存所有层的批量特征或者输入是视频帧序列占用会成倍增长。实际显存占用需要以本机测试为准。控制内存的基本原则是默认保存单张图片的特征。临时保存分析完立即释放。批量图像分析时只保留统计结果不保留全部特征。8.3 降低资源占用的方法减小输入图片尺寸例如从 224x224 降到 112x112特征图面积会变为原来的四分之一。只提取特定层不全量注册 hook。使用output.detach().cpu()尽快把张量从 GPU 移到 CPU释放显存。分析大量图片时每张图片结束后清空feature_maps字典。9. 常见问题与排查方法问题现象可能原因排查方式解决方案预训练权重下载失败网络无法访问下载源查看报错信息确认是超时还是 403配置镜像源或手动下载权重文件放到指定目录hook 回调没有被触发注册的目标层没有被执行打印模型结构确认注册的是实际用到的子模块检查是否把 hook 注册在model顶层而不是具体 layer 上输出特征图全黑特征图范围不在 0~1或全部为负数打印特征图的 min、max 值使用归一化后再绘制绘图报错matplotlib无显示无 GUI 环境或后端不支持尝试matplotlib.use(Agg)保存图片保存为 PNG 文件而不是弹窗显示GPU 显存不足输入图片 batch 过大或保存了太多中间层查看显存占用和报错信息减小 batch及时detach().cpu()使用with torch.no_grad()PyTorch 版本 API 差异不同版本注册钩子的行为有差异检查torch.__version__统一使用register_forward_hook标准用法避免使用早期实验接口特征图分辨率太小看不清深层特征图通常是 7x7 或 14x14查看特征图 shape使用interpolate插值放大到目标尺寸分析结果不稳定模型处于训练模式BatchNorm 行为不同确认model.eval()是否已调用可视化前必须调用model.eval()10. 最佳实践与使用建议10.1 先固定随机种子如果模型包含随机性例如 Dropout 或数据增强需要在脚本开头固定随机种子torch.manual_seed(42) np.random.seed(42) if torch.cuda.is_available(): torch.cuda.manual_seed_all(42)10.2 用小尺寸图片先跑通第一次调试不要直接用高清大图。先用 224x224 或更小的图片跑通流程然后再换成目标分辨率的素材。这样能更快定位是代码问题还是资源问题。10.3 按目录管理输出建议把输入图片、模型权重、输出特征图分目录管理project/ ├── input/ │ └── test.jpg ├── weights/ │ └── resnet18.pth ├── output/ │ ├── feature_maps/ │ └── histograms/ └── scripts/ └── visualize_features.py10.4 使用 hook 收集器封装逻辑如果项目里频繁使用 hook建议封装成一个小工具类方便复用class FeatureExtractor: def __init__(self, model, target_layers): self.model model self.features {} self.handles [] for name, module in model.named_modules(): if name in target_layers: handle module.register_forward_hook( self._make_hook(name) ) self.handles.append(handle) def _make_hook(self, name): def hook_fn(module, input, output): if isinstance(output, torch.Tensor): self.features[name] output.detach().cpu() return hook_fn def __call__(self, x): self.features.clear() with torch.no_grad(): self.model(x) return self.features def remove(self): for handle in self.handles: handle.remove()使用方式extractor FeatureExtractor(model, target_layers[layer2, layer4]) feats extractor(input_tensor)10.5 结合 TensorBoard 扩展除了 matplotlib 静态绘图也可以把特征图写入 TensorBoardfrom torch.utils.tensorboard import SummaryWriter writer SummaryWriter(runs/feat_vis) for name, tensor in feature_maps.items(): writer.add_images(ffeatures/{name}, tensor, dataformatsNCHW) writer.close()这样在训练过程中就可以周期性观察特征变化趋势。11. 总结与下一步建议可视化神经网络中间层输出是 PyTorch 调试技能中性价比非常高的一项。核心只需要掌握register_forward_hook这一套接口就能把任意模型的内部张量导出再用 matplotlib 完成特征图网格、热力图叠加和分布直方图。这篇文章里最值得先跑通的是第三节的 hook 注册代码和第五节的完整可视化脚本。建议先拿一张结构清晰的图片测试观察 conv1、layer2、layer4 的特征差异再对比自己的模型。最容易踩的坑有两个一是忘记model.eval()导致特征分布异常二是深层特征图直接显示时全黑需要先归一化。后续可以继续扩展的方向包括Grad-CAM 类激活图、t-SNE 特征空间降维、特征图相似度分析这些都是在本文 hook 机制基础上的进阶玩法。建议收藏备用实际调试模型时直接按这个流程操作。