如何在ResNet模型的最后分类层添加Softmax激活函数?
在ResNet最后分类层添加Softmax激活函数的方法
有两种简单的实现方式,具体如下:
方法一:修改forward传播函数
直接在forward方法中对模型输出应用Softmax,你可以选择在函数内临时创建Softmax实例,或者在__init__中预先定义模块(后者更高效,避免重复实例化):
方式1.1:临时创建Softmax实例
class ResNet(nn.Module): def __init__(self): super().__init__() self.network = torchvision.models.resnet18() num_ftrs = self.network.fc.in_features self.network.fc = nn.Linear(num_ftrs, 19) def forward(self, xb): logits = self.network(xb) # 对类别维度(dim=1)应用Softmax return nn.Softmax(dim=1)(logits)
方式1.2:在__init__中预定义Softmax模块
class ResNet(nn.Module): def __init__(self): super().__init__() self.network = torchvision.models.resnet18() num_ftrs = self.network.fc.in_features self.network.fc = nn.Linear(num_ftrs, 19) # 预先定义Softmax模块 self.softmax = nn.Softmax(dim=1) def forward(self, xb): logits = self.network(xb) return self.softmax(logits)
方法二:重构全连接层为Sequential组合
把原有的全连接层替换成nn.Sequential容器,将Linear层和Softmax层直接组合在一起:
class ResNet(nn.Module): def __init__(self): super().__init__() self.network = torchvision.models.resnet18() num_ftrs = self.network.fc.in_features # 组合Linear和Softmax作为新的全连接层 self.network.fc = nn.Sequential( nn.Linear(num_ftrs, 19), nn.Softmax(dim=1) ) def forward(self, xb): return self.network(xb)
重要注意事项
- 如果你使用PyTorch内置的
nn.CrossEntropyLoss损失函数,绝对不要额外添加Softmax——该损失函数已经集成了LogSoftmax和NLLLoss的功能,手动添加会导致计算逻辑错误,严重影响模型性能。
内容的提问来源于stack exchange,提问作者seni
相关产品推荐
相关产品推荐

