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

CNN细胞计数模型性能无提升:代码问题排查求助

细胞计数模型验证性能无变化的问题排查

一、数据处理环节的问题

  • 数据增强顺序错误:当前transforms把ToTensor()放在最前面,后续的翻转、旋转、随机裁剪等操作是对张量进行的,而这些增强操作更适合在PIL图像上执行。正确顺序应该是先执行所有PIL层面的增强,再转成张量,最后做归一化:
    transforms.Compose([
        transforms.RandomHorizontalFlip(p=0.5),
        transforms.RandomVerticalFlip(p=0.5),
        transforms.RandomRotation(45),
        transforms.ColorJitter(brightness=0.5, contrast=0.5, saturation=0.5, hue=0.5),
        transforms.RandomResizedCrop(size=image_size, scale=(0.5, 1.0)),
        transforms.ToTensor(),
        transforms.Normalize(mean=[0.485, 0.456, 0.406],
                             std=[0.229, 0.224, 0.225])
    ])
    
  • 冗余的通道重复操作:代码中先将灰度图转为RGB(已得到三通道图像),但后续又把变换后的图像第一个通道重复三次生成新的三通道张量,这会覆盖颜色增强的效果,完全没必要。删除train_resized_image = torch.stack([transformed_train_image[0]] * 3, dim=0),直接将transformed_train_image加入resized_images即可。
  • 验证集预处理缺失:代码仅处理了训练集,未看到验证集的预处理逻辑。验证集应只做固定尺寸调整和归一化,不能使用随机增强,否则会导致验证结果不稳定。需为验证集单独定义无随机操作的transforms,并完成数据加载流程。

二、模型与训练环节的问题

  • 学习率设置过高:Adam优化器的默认推荐学习率是1e-3,而你设置了lr=0.01,过大的学习率会导致模型参数震荡,无法有效收敛。建议调整为lr=1e-4或1e-3,并可根据训练情况添加学习率衰减策略。
  • 模型与数据未移动到GPU:代码中定义了device但未将模型、训练数据移动到对应设备上,导致模型在CPU上训练,不仅速度慢,还可能出现训练与验证设备不一致的问题。修改如下:
    cnn = TotalCellCounter().to(device)
    # 训练循环内
    images, targets = images.to(device), targets.to(device)
    
  • 损失函数维度不匹配:训练时targets的形状是(batch_size,),而模型输出形状是(batch_size,1),虽然PyTorch会自动广播,但可能导致梯度传递异常。需将targets调整为与输出一致的维度:
    loss = criterion(outputs, targets.unsqueeze(1))
    
  • 残差块结构可优化:当前残差块的顺序是Conv->BN->残差相加->ReLU,更标准的残差块顺序是BN->ReLU->Conv,调整后可能提升模型的收敛性:
    # 第一残差块示例
    x1_res = self.conv1_res(x)
    x1 = F.relu(self.bn1(self.conv1(x)))
    x1 = x1 + x1_res
    x1 = F.max_pool2d(x1, kernel_size=2, stride=2)
    

三、验证逻辑的致命错误

  • 错误使用分类任务指标:细胞计数是回归任务,模型输出的是连续的细胞总数,而你用torch.max(outputs.data, 1)计算准确率(分类任务指标)完全不合理,这会导致你看到的“验证性能无变化”是虚假结果。回归任务应使用MAE(平均绝对误差)、RMSE(均方根误差)等指标,修改验证代码如下:
    cnn.eval()
    with torch.no_grad():
        total_mae = 0.0
        total = 0
        for images, labels in valloader:
            images = images.to(device)
            labels = labels.to(device).unsqueeze(1)
            outputs = cnn(images)
            mae = torch.abs(outputs - labels).sum().item()
            total_mae += mae
            total += labels.size(0)
        print(f'验证集平均绝对误差: {total_mae / total:.4f}')
    
  • valloader未定义:代码中直接使用valloader但未看到其创建过程,需确保验证集数据加载流程与训练集一致(使用验证集专属的transforms),避免与训练集混淆。

内容的提问来源于stack exchange,提问作者Ella

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 21:12:52