自定义ResNet101替换avg_pool层遇RuntimeError,迁移学习新手求助
解决ResNet101替换avg_pool层时的Graph Disconnected错误
嘿,作为刚接触迁移学习的新手,遇到这种报错确实头疼!先给你揪出最可能的问题——你写错了池化层的名字!
PyTorch官方提供的ResNet系列(不管是ResNet50还是ResNet101)里,全局平均池化层的名字是avgpool(全小写,没有下划线),而不是你说的avg_pool。你在ResNet50上可能误打误撞用对了操作方式,或者之前修改过层名,但换到ResNet101时,因为层名不匹配,替换操作根本没生效,模型的前向传播路径直接断了,所以才会抛出Graph disconnected的错误。
接下来给你两种靠谱的解决方法,按需选择:
方法一:直接修改原模型(简单快速)
这种方式适合只替换单个层的场景,直接在预训练模型上修改即可:
import torch import torch.nn as nn from torchvision.models import resnet101 # 1. 加载预训练的ResNet101 resnet = resnet101(pretrained=True) # 2. 定义你的自定义池化层 class CustomAvgPool(nn.Module): def __init__(self): super().__init__() # 这里写你的自定义逻辑,比如在池化后加批归一化或dropout self.pool = nn.AdaptiveAvgPool2d((1, 1)) self.bn = nn.BatchNorm2d(2048) # ResNet101的layer4输出通道是2048 def forward(self, x): x = self.pool(x) x = self.bn(x) return x # 3. 替换avgpool层(重点:是avgpool,不是avg_pool!) resnet.avgpool = CustomAvgPool() # 4. 按需修改全连接层(比如适配你的任务类别数) num_classes = 10 # 假设你的任务是10分类 resnet.fc = nn.Linear(2048, num_classes) # 5. 测试前向传播是否正常 test_input = torch.randn(1, 3, 224, 224) output = resnet(test_input) print(output.shape) # 应该输出 torch.Size([1, 10])
方法二:构建新模型类(更清晰,适合复杂修改)
如果之后还要做更多层的调整,推荐用这种方式,能明确控制每一层的连接:
import torch import torch.nn as nn from torchvision.models import resnet101 # 定义自定义池化层 class CustomAvgPool(nn.Module): def __init__(self): super().__init__() self.pool = nn.AdaptiveAvgPool2d((1, 1)) self.dropout = nn.Dropout(0.5) # 举个例子,添加dropout防止过拟合 def forward(self, x): x = self.pool(x) x = self.dropout(x) return x # 构建自定义ResNet101模型 class CustomResNet101(nn.Module): def __init__(self, num_classes=10): super().__init__() # 取出ResNet101的特征提取部分(从conv1到layer4) pretrained_resnet = resnet101(pretrained=True) self.features = nn.Sequential( pretrained_resnet.conv1, pretrained_resnet.bn1, pretrained_resnet.relu, pretrained_resnet.maxpool, pretrained_resnet.layer1, pretrained_resnet.layer2, pretrained_resnet.layer3, pretrained_resnet.layer4 ) # 替换为自定义池化层 self.custom_avg_pool = CustomAvgPool() # 定义新的全连接层 self.fc = nn.Linear(2048, num_classes) # 可选:冻结特征提取层(只训练自定义层和全连接层) for param in self.features.parameters(): param.requires_grad = False def forward(self, x): x = self.features(x) x = self.custom_avg_pool(x) x = torch.flatten(x, 1) # 展平张量后输入全连接层 x = self.fc(x) return x # 测试模型 model = CustomResNet101(num_classes=10) test_input = torch.randn(1, 3, 224, 224) output = model(test_input) print(output.shape) # 输出 torch.Size([1, 10])
额外注意事项
- 替换层后一定要用随机张量测试前向传播,能快速发现连接或维度不匹配的问题。
- 如果还是报错,打印模型结构(
print(model)),检查自定义层是否被正确接入,以及各层的输入输出维度是否匹配。 - 做迁移学习时,按需冻结特征层:如果你的数据集很小,建议冻结大部分特征层,只训练自定义层和全连接层;数据集大的话,可以解冻部分高层特征一起微调。
内容的提问来源于stack exchange,提问作者Rabel Ahmed
相关产品推荐
相关产品推荐

