You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.17 22:53:14