如何将ResNet50改造为可获取中间特征的分层特征提取模型
拆分ResNet50为分层特征提取结构的实现方案
核心思路
ResNet50的特征提取路径天然分为多个阶段,完全可以直接拆分为你需要的分层结构,无需依赖hook:
- 明确ResNet50的特征阶段:初始卷积+池化模块,加上4个残差层(
layer1到layer4),其中layer1输出256通道特征图,layer2输出512通道,layer3输出1024通道,layer4输出2048通道,正好匹配你的需求。 - 剥离分类头:原模型的
fc层是分类专用的,直接忽略即可。 - 重组子模块:将初始模块+
layer1作为第一个特征提取器,后续每个layer单独作为独立子模块,这样就能在forward中依次调用并获取各层输出。
完整实现代码
import torch from torchvision.models import resnet50, ResNet50_Weights import torch.nn as nn class LayeredResNet50(nn.Module): def __init__(self, pretrained=True): super(LayeredResNet50, self).__init__() # 加载预训练ResNet50 resnet = resnet50(weights=ResNet50_Weights.IMAGENET1K_V2 if pretrained else None) # 拆分特征提取模块 # featureExtractor_1: 初始卷积+池化 + layer1,输出256通道特征图 self.featureExtractor_1 = nn.Sequential( resnet.conv1, resnet.bn1, resnet.relu, resnet.maxpool, resnet.layer1 ) # featureExtractor_2: layer2,输入256通道,输出512通道 self.featureExtractor_2 = resnet.layer2 # featureExtractor_3: layer3,输入512通道,输出1024通道 self.featureExtractor_3 = resnet.layer3 # featureExtractor_4: layer4,输入1024通道,输出2048通道 self.featureExtractor_4 = resnet.layer4 def forward(self, x): x1 = self.featureExtractor_1(x) # 输出形状: (batch, 256, H/4, W/4) x2 = self.featureExtractor_2(x1) # 输出形状: (batch, 512, H/8, W/8) x3 = self.featureExtractor_3(x2) # 输出形状: (batch, 1024, H/16, W/16) x4 = self.featureExtractor_4(x3) # 输出形状: (batch, 2048, H/32, W/32) return x1, x2, x3, x4
代码说明
- 模块拆分逻辑:
featureExtractor_1整合了ResNet的初始卷积、归一化、激活、池化以及第一个残差层layer1,处理后得到256通道的特征图,对应你需要的x1。 - 后续子模块:每个
layer直接作为独立模块,因为它们的输入通道数正好匹配前一个模块的输出通道数,无需额外调整。 - 分类头处理:原模型的
fc层属于分类部分,我们不需要,所以初始化时不将其加入自定义模型的子模块即可。
测试验证
# 测试模型输出形状 if __name__ == "__main__": model = LayeredResNet50(pretrained=True) model.eval() # 生成随机输入 (batch_size=2, 3通道, 224x224图像) input_tensor = torch.randn(2, 3, 224, 224) with torch.no_grad(): x1, x2, x3, x4 = model(input_tensor) # 打印各特征图的形状 print(f"x1 shape: {x1.shape}") # 输出: torch.Size([2, 256, 56, 56]) print(f"x2 shape: {x2.shape}") # 输出: torch.Size([2, 512, 28, 28]) print(f"x3 shape: {x3.shape}") # 输出: torch.Size([2, 1024, 14, 14]) print(f"x4 shape: {x4.shape}") # 输出: torch.Size([2, 2048, 7, 7])
内容的提问来源于stack exchange,提问作者dtr43
相关产品推荐
相关产品推荐

