PyTorch预训练ResNet-18特征提取失败问题排查
ResNet-18提取avgpool层512维特征的问题解决与替代方案
为什么forward hook没拿到想要的特征?
你用hook后仍得到fc层输出是因为hook不会修改模型的返回值——模型依旧会执行完整forward流程返回最后一层结果,avgpool的输出需要从hook保存的变量里读取,而不是直接拿模型的返回值。此外还要注意两个常见错误:
- 用不可变类型(比如单个变量)存hook输出:Python中不可变对象在hook函数内赋值无法同步到外部,必须用列表、字典这类可变容器。
- 注册hook后修改了模型对象:比如重新加载模型、覆盖
model变量,会导致之前注册的hook失效。
正确的forward hook用法示例
import torch import torchvision.models as models # 加载预训练ResNet-18并设为评估模式 resnet = models.resnet18(pretrained=True) resnet.eval() # 用列表存储hook输出(列表是可变对象,能在hook内修改) avgpool_features = [] def capture_avgpool(module, input, output): avgpool_features.append(output) # 给avgpool层注册forward hook hook_handle = resnet.avgpool.register_forward_hook(capture_avgpool) # 输入测试图像 test_img = torch.randn(1, 3, 224, 224) # 模型返回的仍是fc层输出 fc_output = resnet(test_img) print("fc层输出形状:", fc_output.shape) # torch.Size([1, 1000]) # 从hook存储的列表中提取avgpool特征,展平为512维 feature = avgpool_features[0].squeeze() print("avgpool层输出形状:", feature.shape) # torch.Size([512]) # 用完hook后移除,避免内存泄漏 hook_handle.remove()
其他提取avgpool特征的方法
方法一:截断模型结构(最常用)
直接去掉ResNet的最后一层fc,让模型输出avgpool的结果:
import torch import torchvision.models as models resnet = models.resnet18(pretrained=True) # 取模型所有子模块,去掉最后一个fc层 feature_extractor = torch.nn.Sequential(*list(resnet.children())[:-1]) feature_extractor.eval() test_img = torch.randn(1, 3, 224, 224) # 输出为[1, 512, 1, 1],squeeze后得到512维向量 feature = feature_extractor(test_img).squeeze() print(feature.shape) # torch.Size([512])
方法二:手动执行forward流程
直接按照ResNet的forward步骤运行,到avgpool层停止:
import torch import torchvision.models as models resnet = models.resnet18(pretrained=True) resnet.eval() test_img = torch.randn(1, 3, 224, 224) x = test_img # 执行ResNet前向流程直到avgpool x = resnet.conv1(x) x = resnet.bn1(x) x = resnet.relu(x) x = resnet.maxpool(x) x = resnet.layer1(x) x = resnet.layer2(x) x = resnet.layer3(x) x = resnet.layer4(x) x = resnet.avgpool(x) feature = x.squeeze() print(feature.shape) # torch.Size([512])
内容的提问来源于stack exchange,提问作者mad
相关产品推荐
相关产品推荐

