使用BCEWithLogitsLoss的ResNet18模型为何仅预测单一类别?
你遇到的核心问题是BCEWithLogitsLoss的使用细节与模型训练流程不匹配,而CrossEntropyLoss因为对应多分类输出模式,恰好避开了这些问题。以下是具体错误点和修正步骤:
错误点分析
1. 模型训练模式未正确开启
如果训练前没有调用model.train(),模型会保持eval模式,梯度不会更新。初始状态下模型的全连接层输出可能普遍偏向负值,导致经过sigmoid后预测结果始终为0。而使用CrossEntropyLoss时你可能无意中开启了训练模式,因此模型能正常收敛。
2. 标签处理冗余且存在隐患
在Dataset中将标签转为torch.long,训练时又重新创建float tensor,不仅冗余,还可能导致设备(CPU/GPU)切换的额外开销,甚至引发隐性错误。
3. 全连接层初始化偏向负输出
BCEWithLogitsLoss依赖logits的正负判断类别,默认的nn.Linear初始化可能让最后一层输出普遍为负,加上如果学习率设置不合理,模型难以快速调整到正确方向。而CrossEntropyLoss对应双输出(num_classes=2),初始化时两个类别的输出更均衡,更容易捕捉类别特征。
4. 类别不平衡未处理(若存在)
如果数据集里0类样本远多于1类,BCEWithLogitsLoss默认不会加权,模型会自然偏向预测多数类;而CrossEntropyLoss在双输出模式下,可能因样本分布或初始化特性,更容易学习到类别差异。
修正步骤
步骤1:确保开启训练模式
在训练循环前添加:
model.train()
步骤2:优化标签处理流程
修改Dataset的__getitem__中标签部分,直接输出匹配模型输出的float类型标签:
# LABELS # label = int(self.target_values[idx]) # 直接转为float32并设置shape为[1,1],与模型输出shape匹配 label = torch.tensor(label, dtype=torch.float32).view(-1, 1) return image, label
同时删除训练代码中冗余的标签转换:
for batch in train_loader: optimizer.zero_grad() inputs, targets = batch inputs, targets = inputs.to(device), targets.to(device) # Forward pass outputs = model(inputs) # 无需再转换targets,已经是float32且shape匹配 loss = loss_fn(outputs, targets) total_loss += loss.item() loss.backward() optimizer.step()
步骤3:调整模型全连接层初始化
为最后一层设置更合理的初始化,避免初始输出偏向负方向:
class ResNet18(nn.Module): def __init__(self, num_classes, band, pt_value): super(ResNet18, self).__init__() resnet = resnet18(pretrained = pt_value) if band != 3: resnet.conv1 = nn.Conv2d(band, 64, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False) self.features = nn.Sequential(*list(resnet.children())[:-2]) self.pool = nn.AdaptiveAvgPool2d(1) self.fc1 = nn.Linear(512, 64) self.fc2 = nn.Linear(64, 1) # 初始化fc2的偏置为0,权重使用He初始化适配relu激活 nn.init.constant_(self.fc2.bias, 0.0) nn.init.kaiming_normal_(self.fc2.weight, mode='fan_in', nonlinearity='relu') def forward(self, x): x = self.features(x) x = self.pool(x) x = x.view(x.size(0), -1) x = relu(self.fc1(x)) x = self.fc2(x) return x
步骤4:处理类别不平衡(可选)
如果数据集0类和1类样本数量差异大,使用带正样本权重的BCEWithLogitsLoss:
# 假设count_0是0类样本数,count_1是1类样本数 pos_weight = torch.tensor([count_0 / count_1], dtype=torch.float32).to(device) loss_fn = nn.BCEWithLogitsLoss(pos_weight=pos_weight)
步骤5:修正预测逻辑
计算准确率时,必须对模型输出应用sigmoid并设置阈值(通常为0.5):
model.eval() correct = 0 total = 0 with torch.no_grad(): for batch in val_loader: inputs, targets = batch inputs, targets = inputs.to(device), targets.to(device) outputs = model(inputs) # 应用sigmoid后判断是否大于阈值0.5 preds = torch.sigmoid(outputs) > 0.5 correct += (preds == targets).sum().item() total += targets.size(0) accuracy = correct / total print(f"Accuracy: {accuracy:.4f}")
内容的提问来源于stack exchange,提问作者Gaia Vallarino

