PyTorch特征提取实战:巧用Hook机制可视化与复用CNN特征
1. 先搞清楚一件事深度学习的特征到底是什么学PyTorch到一定阶段很多人都会卡在一个问题上都说神经网络能自动提取特征那这些特征到底长什么样提取出来又能干嘛我在学完前两篇笔记、跑通几个图像分类的Demo后也在这个问题上纠结了很久。先说个直观的感受。用传统方法做图像特征提取你脑子里想的是颜色直方图、HOG、SIFT这些手工设计的描述子——它们是工程师用数学公式写出来的规则。而深度学习里的特征是网络在大量数据上学出来的中间表示它们没有显式的数学定义藏在每一层卷积输出的张量里。一个224x224的输入图片经过ResNet18的前几层可能变成112x112x64的特征图再过几层变成56x56x128到最后全局池化层出来是一个512维或者2048维的特征向量。这个向量就是对这张图片内容的高度压缩描述。为什么要关心提取特征这件事因为这是迁移学习、图像检索、人脸识别、风格迁移、知识蒸馏等一堆上层应用的地基。比如你想做一个以图搜图的小工具不需要自己从头训练一个分类网络直接用预训练模型把每张图编码成一个特征向量然后算向量之间的余弦相似度就行。再比如你想在很小的数据集上做分类直接从预训练模型里抽出特征去训练一个逻辑回归效果往往好过自己从头训练一个大网络。这篇笔记我打算从实际代码出发覆盖三个层面怎样用PyTorch的hook机制把中间层特征捞出来、怎样把特征图可视化、怎样把最后的特征向量用于下游任务比如相似度检索和简单分类。这些内容是我在自己折腾一个狗品种识别小项目时一步步趟出来的代码都是在CPU和单张GPU上实际跑过的不是抄文档那种干巴巴的示例。2. hook机制从黑盒里把中间特征捞出来的钥匙2.1 为什么不能直接拿中间层的输出如果你有模型推理的完整代码想拿到某一层的输出最直觉的办法是前向传播到那一层的时候手动把张量记录下来。比如你的网络定义是nn.Sequential可以在中间砍一刀class MyNet(nn.Module): def __init__(self): super().__init__() self.features nn.Sequential(...) self.classifier nn.Linear(512, 10) def forward(self, x): feat self.features(x) out self.classifier(feat) return out这样确实能拿到feat但问题是你改了网络的前向逻辑测试时要小心不能影响原始行为如果网络不是你自己定义的而是从torchvision.models里加载的预训练模型比如resnet18(pretrainedTrue)你就得先搞清楚它的内部结构然后继承或改造它很麻烦还容易改错。hook机制就是来解决这个痛点的。它允许你在不修改网络代码的情况下注册一个回调函数在网络前向传播或者反向传播经过某个模块时自动拿到那个模块的输入和输出张量。这就像在电路板上夹一个示波器探针——电路照常工作你只是在旁边观察信号。2.2 从权重文件里反向查结构先看看怎么查网络结构。加载好模型后直接打印import torchvision.models as models model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) print(model)输出会列出完整的模块层次。以ResNet18为例最后几层大概长这样(avgpool): AdaptiveAvgPool2d(output_size(1, 1)) (fc): Linear(in_features512, out_features1000, biasTrue)这里有个很关键的信息我们常说的特征向量在ResNet系列里就是avgpool之后、fc之前那个512维的张量。fc是最后的分类头它输出的1000维数字是每个类别的得分而avgpool之前的特征图带有空间信息池化之后变成1x1x512的向量等价于把整张图的语义信息压缩到了一起。2.3 注册hook的完整代码注册hook很简单我封装了一个小工具类可以在任意层上挂钩子class FeatureExtractor: def __init__(self, model, target_layers): self.model model self.target_layers target_layers self.features {} self._register_hooks() def _register_hooks(self): for name, module in self.model.named_modules(): if name in self.target_layers: module.register_forward_hook(self._make_hook(name)) def _make_hook(self, name): def hook_fn(module, input, output): # output 是一个张量shape 为 (batch_size, channels, H, W) self.features[name] output.detach() return hook_fn def __call__(self, x): self.features {} self.model(x) return self.features用法就是extractor FeatureExtractor(model, target_layers[layer4, avgpool]) feats extractor(input_tensor) # feats[layer4] 是最后一次下采样后的特征图 # feats[avgpool] 是 512 维特征向量但要注意它还是 (batch, 512, 1, 1) 的形状几个细节值得多说一句output.detach()是必须的否则hook会保存整个计算图一次前向还好如果在一个大循环里反复调用显存会爆炸。教训就是这个坑我踩过多跑几百个batch后直接OOM。register_forward_hook返回一个RemovableHandle用完后可以handle.remove()。在写长期运行的服务时一定要清理否则hook会一直累积。从avgpool提取出来的张量形状是(batch, 512, 1, 1)很多新手直接拿去做相似度计算时会报维度不匹配记得先squeeze()或者view(batch, -1)。2.4 不用张量直接看特征图也是可以的如果你想看的是特征图不是向量那hook同样能办到而且更为直观。比如拿layer1的输出shape是(batch, 64, 56, 56)这意味着有64张56x56的图每张图代表了网络在这个尺度上关注的一种模式——可能是边缘、可能是纹理、可能是某种颜色分布。把这些图用torchvision.utils.make_grid拼成一张大图就能很直观地看到网络在看什么。这里给出一个可视化的完整小脚本它可以把指定层的特征图归一化后保存为图片import torch import torchvision.transforms as T from torchvision.utils import make_grid from PIL import Image import matplotlib.pyplot as plt def visualize_feature_maps(feature_map, save_path, nrow8): # feature_map: (C, H, W) C, H, W feature_map.shape # 取前64个通道如果不足就取全部 maps feature_map[:64].unsqueeze(1) # (N, 1, H, W) # 归一化到 [0, 1] maps (maps - maps.min()) / (maps.max() - maps.min() 1e-8) grid make_grid(maps, nrownrow, padding2, normalizeFalse) # grid 是 (1, H, W)转成 PIL 图片 grid_img T.ToPILImage()(grid) grid_img.save(save_path)把这一套用起来你会看到非常有意思的现象浅层特征图很密集像是各种条条框框的边缘检测器深层特征图变得很稀疏大量通道几乎是黑的只有少数几个通道被激活。这说明网络在高层抽象出了一些非常具体的语义概念比如狗耳朵车轮这类部件级别的模式。3. 特征图可视化实操用TensorBoard代替matplotlib3.1 为什么要用TensorBoard很多人提到可视化第一反应是matplotlib。在调试的时候我早期也是用matplotlib把特征图plt.imshow()打出来看但次数一多就烦了——每个epoch、每张图都要手动保存、手动翻文件效率极低。后来改用TensorBoard发现这才是干这个活的正道它可以按标签分组管理所有图片比如layer1/feature_maps、layer2/feature_maps在网页上一键切换。它天然支持标量曲线可以把训练loss、验证acc、特征向量的范数变化一起展示。它对大规模特征图做了切片和缩放优化不会像matplotlib那样一次性把所有子图渲染出来卡到怀疑人生。3.2 一个最小可用的TensorBoard特征图展示方案首先装依赖pip install tensorboardPyTorch 1.2以后自带torch.utils.tensorboard不需要再装tensorboardX。然后在训练或推理循环里把hook收集到的特征图写进去from torch.utils.tensorboard import SummaryWriter writer SummaryWriter(runs/feature_vis) def hook_fn(module, input, output): # output: (batch, C, H, W) x output[0].detach().cpu() # 只取batch里第一张图 C x.shape[0] # TensorBoard 需要 (C, 1, H, W) 或者 (1, C, H, W) 的形式 x x.unsqueeze(1) # (C, 1, H, W) # 归一化 x (x - x.min()) / (x.max() - x.min() 1e-8) writer.add_images(flayer/{name}, x, global_stepstep, dataformatsNCHW) model.layer1[0].register_forward_hook(hook_fn)注意dataformatsNCHW这个参数TensorBoard对图片张量的维度顺序有严格要求默认是NCHW但有时因为前面的squeeze和unsqueeze操作搞乱了就会出图片全是噪点或者通道错乱的怪问题。我记得有一次没传dataformats出来的图片全是红的蓝的条纹排查了半天才发现是维度顺序不对。运行tensorboard --logdir runs浏览器打开http://localhost:6006在Images标签下就能看到每一层的特征图。还能用左上角的滑动条按step翻页对比不同迭代次数下特征的变化。3.3 从看个热闹到看懂门道可视化特征图不是看个花花绿绿就完了。我在实际操作中摸索出一个比较实用的分析套路看稀疏度如果某一层所有特征图都几乎全黑值都接近0说明这个通道在当前输入下没有被激活可能你的输入内容跟预训练数据分布差异很大。看空间结构浅层特征图如果还保留着清晰的空间轮廓说明网络还在编码边缘和纹理深层特征图如果出现大片的模糊高亮区域说明它已经把有没有某个物体部件编码成了一个空间位置。对比不同类别拿猫和狗各一张图分别提取layer4的特征图直观对比哪些通道激活区域不同能帮你理解网络到底在用什么线索做区分。这套分析对于调试自己的网络结构也非常有用。我之前训练一个小的分割模型loss一直不降跑了一遍特征图可视化发现layer1就出现大片饱和区域值全在边缘说明第一层卷积的参数初始化和学习率配合有问题特征直接饱和了。把学习率调低后重新训练损失很快就下去了。如果没有可视化这种问题可能要盲调好几天。4. 把特征向量用起来做相似度检索和简单分类4.1 从图片到向量的一个完整pipeline拿到特征向量之后最直接的应用就是相似度检索。我把自己的狗品种图片库大概3000张图全部过了一遍ResNet18把每张图的512维特征存成一个.npy矩阵然后查一张新图时计算它与所有库内向量的余弦相似度返回top-5。整个pipeline长这样import torch import torchvision.transforms as T import torchvision.models as models import numpy as np from PIL import Image # 1. 定义预处理必须和预训练时保持一致 transform T.Compose([ T.Resize(256), T.CenterCrop(224), T.ToTensor(), T.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # 2. 加载模型去掉fc层直接输出特征向量 model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) model.fc nn.Identity() # 把最后的全连接层替换成恒等映射 model.eval() def extract_feature(img_path): img Image.open(img_path).convert(RGB) x transform(img).unsqueeze(0) # (1, 3, 224, 224) with torch.no_grad(): feat model(x) # (1, 512) feat feat.squeeze() # (512,) feat feat / feat.norm() # 归一化方便后面算余弦 return feat.numpy() # 3. 批量提取库图特征 # features np.stack([extract_feature(p) for p in image_paths]) # 4. 查询 def search(query_path, features_db, paths_db, top_k5): q extract_feature(query_path) scores features_db q # 因为都做了归一化点积就是余弦相似度 top_idx np.argsort(scores)[::-1][:top_k] return [(paths_db[i], scores[i]) for i in top_idx]第2步里model.fc nn.Identity()是关键一行。nn.Identity()是一个恒等层输入是什么输出就是什么而且它不参与梯度计算非常适合用来砍掉分类头。这样前向传播的输出就直接是512维特征向量省掉了hook的麻烦。4.2 做分类时特征提取器能帮你省多少事如果要在自己的小数据集上做分类直接用预训练模型提取特征再训练一个小分类器是非常经典的baseline。以狗品种识别为例假设你只有几百张图从零训练一个ResNet18基本必过拟合但用预训练特征逻辑回归效果往往出奇的好。代码上就两步第一步把所有训练图的特征提取出来存好第二步训练一个sklearn.linear_model.LogisticRegression。from sklearn.linear_model import LogisticRegression from sklearn.model_selection import train_test_split from sklearn.metrics import accuracy_score, classification_report # X: (N, 512) 特征矩阵y: (N,) 类别索引 X_train, X_val, y_train, y_val train_test_split( X, y, test_size0.2, stratifyy, random_state42 ) clf LogisticRegression(max_iter1000, C1.0) clf.fit(X_train, y_train) y_pred clf.predict(X_val) print(accuracy_score(y_val, y_pred))我自己的小项目用这套方案几百张训练图就达到了大约85%的top-1准确率而用同样数据集从头训ResNet18只能到60%多点。如果后面把逻辑回归换成MLP或者继续微调整个网络还能往上走。这种工作流最大的好处是快——提取特征和训练分类器都在分钟级别方便你快速验证这个任务在预训练特征上到底有多少可分性。4.3 特征向量的语义继承为什么它能直接拿来用很多人会问一个很本质的问题预训练模型是在ImageNet上训的凭什么它提取的特征能用在狗品种这种完全不同的任务上答案是ImageNet的1000类里本身就包含大量狗品种大约120多种所以ResNet在狗这个语义区域上已经学到了丰富的判别特征。退一步说即使你的任务跟ImageNet完全不搭边比如识别医疗影像、检测工业缺陷卷积网络的前几层提取的边缘、纹理、颜色过渡这类低级特征是视觉任务的通用基础迁移到新任务上依然有效。越是深层特征越跟源任务绑定越是浅层特征越通用。所以做迁移的时候如果你数据量极小就只训分类头数据量稍多再放开后面一两个block微调这是一个基本策略。5. 把模型当特征提取器用留意预处理和mode的坑5.1 预处理不一致特征就白提了这是所有人都会踩、但很少人提前告诉你的坑。你加载的预训练模型是在特定预处理下训练的。torchvision官方模型的默认预处理是Resize(256)-CenterCrop(224)-ToTensor()-Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225])。这些数字是ImageNet数据集的RGB均值和标准差不是随便拍的。如果你自己写了个transform忘了归一化或者用了别的均值方差那么喂进去的图像分布跟预训练时完全不同前面的卷积层提取的特征会严重失真。这里特别提醒很多人会不小心做了两次归一化比如在数据预处理脚本里归一了一次在提取特征时又归一了一次特征向量看起来还是有数值但语义信息已经不对了。我建议把预处理封装成一个全局函数所有提取特征的地方都调用同一个避免各处写一份造成不一致。5.2 model.eval() 不是可选项PyTorch的nn.Module默认是train模式在train模式下BatchNorm会用当前batch的统计量做归一化Dropout层会随机丢弃神经元。如果你提取特征时忘了切到eval()模式BatchNorm的计算结果是不正确的——它用的不是训练时积累的全局统计量而是当前这个batch的统计量。小batch比如1张图时这个问题尤其严重单张图片的均值和方差噪声很大导致特征向量不稳定。同一张图不同batch大小提取出来的特征甚至差异巨大。我第一次做检索时始终找不到为什么同一张库图每次query结果不一样最后发现问题就出在这一行代码上。model.eval() # 必须调用 with torch.no_grad(): feat model(x)5.3 显存不够时用CPU推理做特征提取完全可行特征提取这个场景推理时batch size不需要很大。一次处理一张图ResNet18在CPU上大概需要50~100ms3000张图也就几分钟的事。如果你的机器有GPU但显存不大完全没必要硬上大batch反而是batch_size1配合torch.no_grad()更稳。另外如果你是为了提取特征去训练分类器还可以多提一个数据增强版本对同一张图做随机裁剪和水平翻转提取多个增强视图的特征取平均即测试时增强TTA。这个方法可以显著提升特征的质量。我之前用了一版TTA检索精度大概涨了2~3个点。6. 特征可视化与下游任务之外的延伸方向6.1 用t-SNE验证特征质量特征提出来以后怎么评价它的质量除了做下游任务看指标还有一个非常直观的方法把特征降维到2维或者3维画散点图看看不同类别的样本是不是自然聚成团。sklearn.manifold.TSNE是绕不开的工具。把512维特征降到2维画出来是这样的逻辑from sklearn.manifold import TSNE import matplotlib.pyplot as plt feats_tsne TSNE(n_components2, perplexity30, initpca, random_state42).fit_transform(features) plt.figure(figsize(8, 8)) scatter plt.scatter(feats_tsne[:, 0], feats_tsne[:, 1], cy, cmaptab10, s8) plt.colorbar(scatter) plt.savefig(tsne_feats.png, dpi200)这里有个参数需要提一下perplexity可以理解为每个点周围有效邻居的数量。一般取值在5到50之间数据量小就不妨用小的perplexity数据量大可以调大一点。我在3000张图、30个类别时用30效果还不错。看t-SNE图的时候有个经验如果同一类别的点散成一团、跟其他类别完全重叠说明你提取的特征在这个任务上没有区分度这时候别急着调分类器回头看看是不是预处理出了问题、或者特征提取的层选得不对。反过来如果类间边界清晰、类内聚集紧凑下游任务基本稳了。6.2 从单层到多层把特征金字塔用起来单一层特征有时候不够用尤其是任务里有大大小小不同尺度的目标。这时候可以提取多个层级的特征拼接起来形成所谓的特征金字塔。比如同时提取layer2输出56x56、layer3输出28x28、layer4输出14x14的特征图分别做全局池化拼成一个更长的特征向量。代码上不需要复杂改动还是那套hook只是多加几个目标层然后在拼接时注意维度对齐feat_l2 torch.nn.functional.adaptive_avg_pool2d(feats[layer2], 1).squeeze() feat_l3 torch.nn.functional.adaptive_avg_pool2d(feats[layer3], 1).squeeze() feat_l4 torch.nn.functional.adaptive_avg_pool2d(feats[layer4], 1).squeeze() multi_scale_feat torch.cat([feat_l2, feat_l3, feat_l4], dim0) # 长度 128256512多尺度特征对检索和分类都有帮助代价是特征维度变大、计算量变多但对几万张图的规模来说完全不是问题。6.3 特征提取与finetune怎么衔接最后聊一下什么时候该用直接提取特征、什么时候该微调模型。我的经验是当作分类任务时先用提取特征逻辑回归快速跑一个baseline如果baseline的准确率能接受就不用微调省时省力如果baseline离你的目标还差很多再考虑用预训练权重初始化然后整体或者部分层在你的数据集上微调。微调的起步建议是先用较小的学习率比如1e-4量级训练所有层但把backbone的requires_grad设为False只训练分类头等loss稳定后再解开backbone后几层的梯度做端到端微调。这样既避免了一上来就破坏预训练特征也给了模型充分时间适应新数据。7. 这些坑我替你踩过了坑一hook不清理显存被吃光如果hook注册次数过多比如在训练循环里反复注册且output没有detach()计算图会一直挂着显存只增不减。一定要把hook的handle保存下来不用了马上移除。如果你只是推理甚至不需要hook直接用model.fc nn.Identity()更简洁。坑二TensorBoard的图片格式问题add_images要求输入的shape是(N, C, H, W)并且值域最好是[0, 1]浮点型或者直接传无符号整型图。特征图的取值范围是没有约束的一定要自己归一化否则出来的图要么全黑要么全白。坑三BatchNorm在train模式下的假特征前面提过的老问题再强调一次提取特征前必须先model.eval()。很多人在写了model.eval()之后又调用了model.train()或者忘记让no_grad()包裹推理都会造成特征不一致。检查方式很朴素同一张图跑两次对比特征向量是否完全相同如果不同赶紧查。坑四CNN输入分辨率不要太随意预训练模型对输入分辨率是有心理预期的。你用Resize(256)CenterCrop(224)模型叹口气说还行你要是直接Resize(448)某些模型会因为感受野变化导致性能下降。ResNet结构本身对尺寸不敏感全局池化层统一了维度但特征的语义质量会受影响。想稳妥就按官方推荐尺寸来。坑五特征向量要不要做归一化做相似度检索时强烈建议对特征向量做L2归一化这样可以只算点积就能得到余弦相似度速度更快。但如果是用于训练逻辑回归归一化与否影响不大。我习惯把归一化放在提取阶段就做了这样下游用起来很干净。所以总的来说特征提取这条链路的核心就三件事选对层、用对模式、做对预处理。只要这三步不走偏不管你是做可视化、检索还是做下游分类都能省下大量的重复劳动。现在再回头看PyTorch的模型不再是一个只能吃图片吐分类概率的黑盒而是可以随时打开检查的透明工具了——这种感觉比跑通任何一个教程Demo都爽得多。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →