
这几年做图像模型我几乎每次调试都离不开特征图提取和特征图可视化这两件事。Pytorch 生态里实现这个操作绕不开 hook 机制但网上很多教程把 hook 讲得过于复杂动辄几十行封装成工具类。其实最核心的流程六行代码完全能跑通。这篇文章就从这六行代码展开把特征图提取和特征图可视化的原理、细节、坑点一次讲清楚适合想深入理解卷积神经网络到底在“看什么”的朋友也适合正在做模型调试、论文配图、缺陷检测可视化的同学直接抄作业。1. 特征图可视化是什么调试模型时能帮你看到什么1.1 一张特征图背后模型在“看”什么卷积神经网络本质上是一个逐层抽象的特征提取器。第一层卷积可能在高频细节上响应比如边缘、角点、颜色渐变中间的卷积层开始组合出纹理、局部形状到了最后的卷积层每个通道已经能对应到某种语义部件比如眼睛、轮子、窗户这类概念。这些中间结果就是特征图feature map它的形状通常是 [B, C, H, W]代表批大小、通道数、特征图高度和宽度。很多初学者会把 CNN 当作一个“黑盒”输入一张图输出一个类别概率中间过程完全看不到。特征图可视化就是把这个黑盒打开一条缝让你亲眼看到每一层提取到了什么信息。我自己的体会是它是排查模型“没学好”还是“数据有问题”最直接的证据。比如做一个缺陷检测模型训练损失一直降不下去如果可视化最后一层特征图发现它根本没有区分出缺陷区域那问题就不在分类头而是前面特征提取根本没有学到有效信息这时候再调学习率、换损失函数都是白费力气。另一个常见场景是论文配图。投稿时审稿人特别爱看特征图可视化因为几张图就能说明你的网络结构确实在提取有意义的特征而不是单纯地过拟合训练集。做模型压缩、蒸馏、可解释性研究时特征图可视化更是绕不开的基础工具。1.2 什么场景下必须用到特征图可视化我把日常需要特征图可视化的场景分成三类你在实际项目中可以对号入座模型调试判断某一层是否学到有效特征、特征是否有冗余、是否发生了梯度消失导致的“死层”整个特征图全零或噪声。结构优化对比不同网络结构比如 ResNet 和 MobileNet在同一层上的特征差异决定是否上注意力模块、是否增加通道数。结果展示在项目汇报、论文、技术博客里用可视化图直观体现方法的有效性。尤其是第三类它几乎成了深度学习论文的标配。之前我帮一个师弟调目标检测模型他把最后一层特征图导出来叠加到原图上一眼就看出小目标区域响应弱到几乎看不见。这个结论如果只看 mAP 指标需要做好几组消融实验才能发现。2. 六行核心代码hook机制的极简实现2.1 为什么要用 hook 来提取特征图提取特征图最直接的想法是修改模型 forward 函数在中间层 return 出来。但实际工程里模型往往来自 torchvision 或者第三方库直接改源码风险很大万一 forward 内部还有其他逻辑比如多分支、特征融合稍不留神就改坏了。Pytorch 官方早就想到了这个问题提供了register_forward_hook机制。这个名字听着高大上其实理解起来很简单你给某个模块装了一个“监听器”每次这个模块前向传播执行完监听器就会自动被调用你能在此时拿到它的输入和输出。整个过程完全不需要改动原始网络结构用完也可以随时移除。对比一下两种方案方案改动程度风险灵活性修改 forward 源码高可能破坏原有结构每次都要改代码注册 hook零侵入低可随时注册/移除可插拔实际使用中hook 还有一个好处可以同时挂到很多层上一次性拿到多个中间层的输出。比如我想同时看 ResNet18 的 layer1、layer2、layer3、layer4 四个层只需要循环注册四个 hook 就行不需要改任何 forward 代码。2.2 六行代码逐行拆解直接上代码这六行是核心骨架# 1. 初始化一个字典用来存特征图 features {} # 2. 定义一个 hook 函数把输出存进字典 def hook_fn(module, input, output): features[layer4] output.detach() # 3. 注册 hook 到目标层 model.layer4[2].conv2.register_forward_hook(hook_fn) # 4. 切到 eval 模式关闭梯度计算 model.eval() with torch.no_grad(): out model(input_tensor) # 5. 从字典里取出特征图 feature_map features[layer4] # 6. 用 torchvision 自带的工具保存拼图 save_image(make_grid(feature_map[0].unsqueeze(1), nrow8, normalizeTrue), feature_map.png)我逐行解释一下你可能疑惑的点。output.detach()这行非常关键。特征图在 forward 过程中是带着完整计算图的因为我们只想做可视化不需要反传梯度。如果不 detach整个计算图会一直挂在features字典里等这个变量被覆盖或者程序结束内存才被释放。在循环推理的场景下这会造成内存持续增长跑几百张图就 OOM。detach 之后feature_map就是一个独立张量和计算图彻底断开了。model.layer4[2].conv2是注册位置。ResNet18 的 layer4 里有 2 个 BasicBlocklayer4[2]是最后一个 BasicBlock.conv2是它内部第二次卷积。选这一层是因为它输出的是整个骨干网络最后的特征图语义信息最丰富。如果你用别的模型这一步需要换成目标模型的实际层名。第 6 行的make_grid和save_image是 torchvision 的实用函数。feature_map[0]取第一个样本的特征图形状是 [512, 7, 7]unsqueeze(1)把它变成 [512, 1, 7, 7]这是 make_grid 要求的输入格式它会把这 512 个单通道小图排列成一张大图网格。nrow8表示每行放 8 个小图normalizeTrue表示自动归一化到 [0,1] 区间再保存。2.3 核心注意事项detach、eval 与 hook 的生命周期刚才说了 detach 的重要性这里再补充两个容易被忽略的点。第一model.eval()一定要加。PyTorch 里 BN 层和 Dropout 层在训练和推理状态下的行为完全不同。如果不切到 eval 模式BN 层会使用当前 batch 的统计量而不是全局统计量导致特征图不稳定每次前向结果都不一样。而且 hook 里拿到的 output 会带上 BN 层的中间状态可视化出来可能充满噪声。我的习惯是任何只做推理不做训练的场景第一件事就是model.eval()。第二hook 会在每次前向传播时都被调用。如果你的代码里多次调用model(input_tensor)features字典里的值会被反复覆盖。在测试集上批量提取特征图时这一点反而很方便每张图前向一次然后立刻把特征图拿去保存最后字典里留下的是最后一张图的特征图。但如果你把 hook 留在模型上一直不清除后面跑反向传播训练时这个 hook 同样会触发白白增加开销。所以在训练脚本里用完 hook 最好调用handle.remove()移除。handle model.layer4[2].conv2.register_forward_hook(hook_fn) # 用完移除 handle.remove()3. 完整实操流程从单张图片输入到可视化图片落地3.1 环境准备与输入预处理六行代码能跑通是在理想条件下实际做特征图可视化还需要把数据预处理、通道顺序、图片尺寸这些细节处理干净。我用最常见的 ResNet18 加一张自然图片做演示。环境方面安装好torch和torchvision就行CPU 也能跑但用 GPU 会快很多。如果还没装直接pip install torch torchvision或者用 conda 创建环境安装这块不同系统的教程很多不展开了。输入预处理有几件事必须做。torchvision 的预训练模型默认是在 ImageNet 上训练的输入的图片会被归一化到某个特定分布。你需要用 ImageNet 的均值和标准差做标准化而且顺序不能反先 resize、再转 tensor、再 normalize。from torchvision import transforms from PIL import Image 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]) ]) img Image.open(demo.jpg).convert(RGB) input_tensor transform(img).unsqueeze(0) # [1, 3, 224, 224]unsqueeze(0)很关键模型要求输入是四维的 [B, C, H, W]而单张图片处理后是三维 [C, H, W]必须补上 batch 维度否则模型会直接报错。这里有个小细节Resize((224, 224))是直接拉伸图片而不是等比例缩放如果原图本身不是方形拉伸会改变物体比例对分类模型影响不大但如果你想做精确的热力图叠加最好用CenterCrop配合Resize保持比例。更稳妥的做法是先缩放到短边 256再中心裁剪到 224这是 ImageNet 标准流程transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ])3.2 特征图拼图输出make_grid 的正确用法拿到特征图 [1, 512, 7, 7] 之后很多人直接save_image(feature_map, out.png)然后发现保存出来的图是黑的或者提示 channel 不对。原因在于save_image默认要求输入是 [B, C, H, W] 并且值在 [0,1] 区间。特征图里的值可能是负数也可能超过 1直接保存肯定不对。正确做法是先处理维度和值域。make_grid接收 [N, C, H, W] 的输入对于灰度特征图我们要把通道维度显式表示为 1所以feature_map[0].unsqueeze(1)让 [512, 7, 7] 变成 [512, 1, 7, 7]。然后设置normalizeTruemake_grid会计算整个网格里的最小值和最大值然后线性拉伸到 [0,1] 区间。不过normalizeTrue是全局归一化如果某几个通道的激活值特别大其他通道会被压缩得很暗细节看不清。更精细的做法是手动per-channel归一化用torch.nn.functional.normalize或者自己算def normalize_feature_map(fm): # fm: [C, H, W] fm_min fm.amin(dim(1, 2), keepdimTrue) fm_max fm.amax(dim(1, 2), keepdimTrue) return (fm - fm_min) / (fm_max - fm_min 1e-6) normalized normalize_feature_map(feature_map[0]) grid make_grid(normalized.unsqueeze(1), nrow8, padding2) save_image(grid, feature_map_grid.png)实际操作里我会先把 512 个通道按激活强度排序把响应最强烈的前几十个通道单独拿出来看这样能快速定位模型主要关注什么。如果直接看 512 张小图拼出来的大图人眼很难分辨出哪些通道更重要。3.3 特征图与原图叠加热力图可视化特征图网格图能看出通道间的差异但很难直观反映“这个区域被模型关注了”。更直观的方式是把特征图转成热力图叠加到原图上。这也是论文里最常见的展示形式。我们需要把 [512, 7, 7] 的特征图压缩成单通道的显著性图。通常做法是所有通道求平均也可以取最大值得到 [7, 7] 的矩阵再 resize 到和原图一样大用 cv2 的伪彩色映射转换成彩色图最后与原图加权融合。import cv2 import numpy as np import torch.nn.functional as F # 1. 特征图压缩成单通道 saliency feature_map[0].mean(dim0) # [7, 7] saliency F.relu(saliency) # 只保留正响应 # 2. resize 到原图尺寸 saliency_np saliency.cpu().numpy() saliency_resized cv2.resize(saliency_np, (img_width, img_height)) # 3. 归一化到 0-255 saliency_norm cv2.normalize(saliency_resized, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8) # 4. 伪彩色映射 heatmap cv2.applyColorMap(saliency_norm, cv2.COLORMAP_JET) # 5. 与原图融合 original cv2.cvtColor(np.array(img), cv2.COLOR_RGB2BGR) overlay cv2.addWeighted(original, 0.5, heatmap, 0.5, 0)融合时 alpha 的选择影响很大0.5 是中间值适合大部分场景。如果你想强调热力图区域可以把热力图权重调大到 0.7想看清楚原图细节就把原图权重调大。我个人的经验是在演示时用 0.5 比较均衡在分析具体缺陷位置时用 0.3 的热力图权重可以更清晰地看出响应区域对应原图的哪个部分。这里踩过一个坑cv2 读图是 BGR 通道顺序而 PIL 读出来是 RGB。如果直接用 PIL 转 numpy 再交给 cv2 处理颜色会明显偏蓝。所以上面代码里我加了一步cvtColor(img, COLOR_RGB2BGR)这个细节不处理叠加出来的热力图颜色看起来就是不对的。4. 进阶玩法多层级特征对比与类激活图4.1 多层级特征对比从边缘纹理到语义部件只看最后一层特征图其实丢失了很多信息。我强烈建议你在一开始就把模型的多个层都挂上 hook同时提取不同层级的特征图看它们在学什么。以 ResNet18 为例layer1 输出的特征图是 [256, 56, 56]保存网格图时单张小图只有 56x56 像素勉强能看出边缘和纹理的响应layer2 是 [256, 28, 28]layer3 是 [512, 14, 14]layer4 是 [512, 7, 7]。从低层到高层你会看到特征图的小图越来越“抽象”空间分辨率越来越低但每个通道的语义越来越明确。具体实现上只需要一个循环注册多个 hookfeature_maps {} def make_hook(name): def hook_fn(module, input, output): feature_maps[name] output.detach() return hook_fn # 注册多个层 model.layer1[-1].register_forward_hook(make_hook(layer1)) model.layer2[-1].register_forward_hook(make_hook(layer2)) model.layer3[-1].register_forward_hook(make_hook(layer3)) model.layer4[-1].register_forward_hook(make_hook(layer4))注意make_hook函数构建了一个闭包目的是让每个 hook 记住自己的层名。如果不这样写直接在循环里用同一个name变量最后所有 hook 拿到的都是最后一个名字这是 Python 经典闭包陷阱很多人刚写时都中招过。多层对比最有价值的场景是模型诊断。有一次我调一个细粒度分类模型发现 layer1 和 layer2 的特征图激活很丰富但 layer3 和 layer4 的特征图大量通道都是近乎全零的噪声。这个现象说明特征提取到深层后信息流失严重多半是梯度传播出了问题。顺着这个线索我把目光锁定在 BatchNorm 层和残差连接上最后发现是初始化方式导致深层梯度不稳。这种问题如果不做特征图可视化可能要调一周的参数才能偶然发现。4.2 Grad-CAM 和特征图可视化的关系特征图可视化是直接看中间层的输出而 Grad-CAM梯度加权类激活映射是在特征图的基础上用类别得分对特征图的梯度作为权重对通道进行加权求和得到一张类别相关的热力图。两者最大的区别是特征图可视化展示的是模型在该层“提取到什么”Grad-CAM 展示的是模型“因为什么区域而做出判断”。举个例子一张猫的照片输入模型某张特征图可能在猫的胡须区域有强响应另一张可能在耳朵区域有强响应这些是模型提取到的特征而 Grad-CAM 会告诉你对于“猫”这个类别模型主要看的是耳朵和胡须区域而不是背景的沙发。Grad-CAM 的计算核心可以用下面这段代码理解def grad_cam(model, input_tensor, target_class): features None grads None def forward_hook(module, input, output): nonlocal features features output.detach() def backward_hook(module, grad_input, grad_output): nonlocal grads grads grad_output[0].detach() handle_f model.layer4[-1].register_forward_hook(forward_hook) handle_b model.layer4[-1].register_full_backward_hook(backward_hook) output model(input_tensor) model.zero_grad() one_hot torch.zeros_like(output) one_hot[0][target_class] 1 output.backward(gradientone_hot) weights grads.mean(dim(2, 3), keepdimTrue) # GAP cam (weights * features).sum(dim1, keepdimTrue) cam F.relu(cam) handle_f.remove() handle_b.remove() return camGrad-CAM 的输出是一张和特征图尺寸相同的热力图后续的 resize、归一化、叠加到原图流程和前面介绍的热力图完全一样。我把 Grad-CAM 和特征图可视化搭配使用先用特征图可视化判断“模型提取到了什么”再用 Grad-CAM 判断“模型重点关注哪里”两者互为补充基本可以覆盖大多数模型解释需求。4.3 中间层特征降维可视化PCA 与 t-SNE这里再分享一下我用过的另一种思路。当特征图的通道数非常多比如上千通道时直接看网格图几乎没有意义密密麻麻的小图让人眼花缭乱。这时可以把特征图当作高维特征向量用 PCA 降到 3 通道以伪彩色图的形式展示主要成分。具体来说将 [C, H, W] 的特征图变形为 [C, H*W]用 PCA 降维到 3 维再重新排列回 [3, H, W]三通道分别对应三个主成分。这样得到的图能呈现出特征空间的整体分布结构。t-SNE 则常用于不同样本之间的特征分布对比它会将每张图像的特征向量映射到二维平面看同类样本是否聚在一起。这两个方法配合特征图可视化可以更立体地理解模型行为。5. 常见问题与排查技巧实录5.1 hook 没有被调用特征图是空的这类问题最典型。代码逻辑看起来都对features字典却一直是空的或者保存时提示 KeyError。排查思路按顺序来第一步确认注册的模块真的在模型前向传播中被执行了。比如有人给model.fc注册 hook但 ResNet 的fc在骨干网络之后如果输入数据没有经过model.forward或者是被model.features单独调用了fc的 hook 不会被触发。最简单的验证方法是在 hook 函数第一行加print(hook called)如果打印都没出现说明路径根本没走到。第二步确认模型不是被torch.compile或类似封装包过。PyTorch 2.0 之后torch.compile会把模型图优化某些情况下 hook 的行为会受影响。如果项目用了 compile 导致 hook 失效一个临时规避办法是改用原始模型推理或者在 compile 之前注册 hook。第三步检查模块名是否存在。model.layer4[2].conv2需要在模型结构里真实存在。可以用for name, _ in model.named_modules(): print(name)打印出所有模块名然后一个一个对照。5.2 图像全黑、全白、颜色不对特征图可视化出来的图不对绝大多数是值域和通道顺序问题。全黑通常是特征图的值全部在 0 附近或者是负值大于正值导致 normalize 后大部分区域接近 0。可以打印feature_map.min()和feature_map.max()确认数值范围。全白往往是某些 batchnorm 层输出偏移太大特征图的值整体很大归一化后被压成一片白。颜色不对偏蓝、偏绿几乎都是 RGB/BGR 顺序问题。前面我提过 cv2 和 PIL 的颜色通道顺序不同记住一个原则如果是用 PIL 读图最后给 cv2 时先转一次COLOR_RGB2BGR如果全程用 cv2 读图就不用转。另一个容易忽略的是save_image保存后如果你用图片查看器打开看到的颜色也可能因为软件本身的色彩管理而有偏差。最好用 PIL 或者 matplotlib 重新读一遍存的文件确认颜色是正常的再拿去用。5.3 显存与内存爆炸特征图可视化一般都配合torch.no_grad()使用但即便关闭了梯度大量地保存特征图依然会占显存。比如输入 [16, 3, 224, 224] 的 batch一个 [16, 512, 7, 7] 的特征图虽然本身不大但如果同时在多个层保存以及多轮推理没有及时清空字典累积起来的显存占用就很可观。处理方案有几个在 hook 里保存时就转成 CPUoutput.detach().cpu()这样显存里不保留特征图只保留一份 CPU 副本。及时清理字典用完一个特征图就del features[layer4]并torch.cuda.empty_cache()注意这并不能完全释放显存只是让显存碎片可以被后续分配复用。如果只是可视化不需要很多通道可以只保留激活最强的前 K 个通道。在 hook 里直接取前 K 个通道再存能大幅减小存储压力。我之前在服务器上跑一个大规模特征图提取任务输入是视频帧序列刚开始没注意每帧都保存整套特征图跑了半小时显存直接爆掉。后来改成在 hook 里就压缩、转 CPU、只保留关键通道显存占用缩小到原来的十分之一。5.4 模型输出与预期类别不一致做可视化时一个隐含假设是模型确实学到了你期望的特征。但如果你拿一张猫的图片模型却预测成狗特征图可视化会变得很奇怪因为模型提取的特征分布是基于“狗”这个类别的决策路径。所以我在调试特征图前通常先打印一下model(input_tensor)的输出概率确认预测类别正确。如果类别都不对优先排查图像的预处理是否和训练时一致是否有选错模型权重图像内容是否太模糊或太小导致模型识别不出来这些小问题不先确认直接可视化特征图很容易误判成模型结构有问题。6. 一个直接能跑的完整示例脚本把前面的代码整合成一个可以直接跑通的脚本方便你保存下来当模板使用。import torch import torch.nn.functional as F import torchvision.models as models from torchvision import transforms from torchvision.utils import make_grid, save_image from PIL import Image import cv2 import numpy as np # 1. 加载模型 device torch.device(cuda if torch.cuda.is_available() else cpu) model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) model.to(device) model.eval() # 2. 定义 hook feature_maps {} def make_hook(name): def hook_fn(module, input, output): feature_maps[name] output.detach().cpu() return hook_fn model.layer4[-1].register_forward_hook(make_hook(layer4)) # 3. 图像预处理 transform transforms.Compose([ transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) img Image.open(demo.jpg).convert(RGB) input_tensor transform(img).unsqueeze(0).to(device) # 4. 前向传播 with torch.no_grad(): output model(input_tensor) print(预测类别:, output.argmax(dim1).item()) # 5. 保存特征图网格 fm feature_maps[layer4][0] # [C, H, W] fm F.relu(fm) fm_norm (fm - fm.amin(dim(1,2), keepdimTrue)) / \ (fm.amax(dim(1,2), keepdimTrue) - fm.amin(dim(1,2), keepdimTrue) 1e-6) grid make_grid(fm_norm.unsqueeze(1), nrow8, padding2) save_image(grid, feature_map_grid.png) # 6. 保存叠加热力图 saliency fm.mean(dim0).numpy() saliency cv2.resize(saliency, (img.width, img.height)) saliency cv2.normalize(saliency, None, 0, 255, cv2.NORM_MINMAX).astype(np.uint8) heatmap cv2.applyColorMap(saliency, cv2.COLORMAP_JET) original cv2.cvtColor(np.array(img), cv2.COLOR_RGB2BGR) overlay cv2.addWeighted(original, 0.5, heatmap, 0.5, 0) cv2.imwrite(overlay.png, overlay)这个脚本我建议你直接跑一遍然后换成自己的图片、自己的模型观察不同层特征图的差异。跑通了之后再往里面加自己的逻辑比从零开始拼代码效率高很多。最后再分享一个我自己的使用习惯。特征图可视化最容易犯的错误是只保存一张图拿过来看一眼然后就没有然后了。我一般会构建一个简易的“可视化矩阵”横轴是不同输入样本比如正常样本、异常样本、攻击样本纵轴是不同层或不同通道这样能快速对比同一层在不同输入上的响应差异。矩阵式可视化在排查数据集 bias 时特别好用当两个类别的特征图激活区域几乎完全重合时说明模型没有学到判别性特征该换网络结构或者加注意力模块了。特征图提取这个概念本身很简单但能玩到什么深度取决于你愿不愿意花时间去逐层观察、逐个样本对比。