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

segmentation_models_pytorch分割训练报AssertionError修复方法

问题根因

报错本质是代码中存在大量mask/masks变量名拼写错误,传入DiceLoss的不是当前批次的真实掩码张量,导致批次维度不匹配触发断言,同时存在多处隐藏bug会阻塞后续运行:

  • 模型forward方法定义入参为masks=None,但损失计算分支判断条件错误写为if mask != None,引用了不存在的变量
  • 训练、验证循环中从DataLoader取出批次数据为images, masks,但设备迁移代码错误写为masks = mask.to(DEVICE),错误引用未定义的单数形式mask,传入损失的是内存中残留的无关张量,批次维度自然和模型输出不匹配
  • 验证阶段数据增强错误使用RandomCrop随机裁剪,会导致验证结果波动、输出尺寸不稳定
  • 训练函数内部缩进错误,train_loss = train_func(...)被写在函数体内部永远无法执行
  • 验证集DataLoader的batch_size被设置为整个验证集样本数,样本量稍大就会触发显存溢出
  • 每次模型前向传播都重新实例化DiceLoss,存在不必要的性能开销
修复方案

按以下步骤修改对应代码即可解决问题:

  1. 修正模型类,将损失函数初始化移到构造函数,统一变量名:
class SegmentationModel(nn.Module):
  def __init__(self):
    super(SegmentationModel,self).__init__()
    self.lossF = DiceLoss(mode='binary')
    self.backbone = smp.Unet(
        encoder_name=ENCODER,
        encoder_weights=WEIGHTS,
        in_channels=3,
        classes=1,
        activation=None
      )
  def forward(self,images, masks=None):
    logits = self.backbone(images)
    if masks is not None:
      return logits, self.lossF(logits,masks)
    return logits
  1. 修正验证集增强,替换随机裁剪为固定尺寸缩放:
def get_val_augs():
  return A.Compose([
    A.Resize(width=IMG_SIZE, height=IMG_SIZE),
  ])
  1. 修正训练函数的变量名错误,移除函数内部错误缩进的调用代码:
def train_func(dataloader, model,optimizer):
  model.train()
  total_loss = 0.0
  for images, masks in tqdm(dataloader):
    images = images.to(DEVICE)
    masks = masks.to(DEVICE)

    optimizer.zero_grad()
    logits, loss = model(images,masks)
    loss.backward()
    optimizer.step()
    total_loss += loss.item()
  return total_loss / len(dataloader)
  1. 修正验证函数的变量名错误:
def eval_func(dataloader, model):
  model.eval()
  total_loss = 0.0
  with torch.no_grad():
    for images, masks in tqdm(dataloader):
      images = images.to(DEVICE)
      masks = masks.to(DEVICE)
      logits, loss = model(images,masks)
      total_loss += loss.item()
    return total_loss / len(dataloader)
  1. 修正数据集初始化、DataLoader配置和训练循环逻辑,不要把验证集batch_size设为全量样本:
# 初始化数据集
trainset = SegmentationDataset(Train_DF, get_train_augs())
valset = SegmentationDataset(Val_DF, get_val_augs())

# 配置DataLoader
trainloader = DataLoader(trainset, batch_size=BATCH_SIZE, shuffle=True)
# 验证集batch_size和训练集保持一致即可,不要设为len(valset)
validloader = DataLoader(valset, batch_size=BATCH_SIZE, shuffle=False)

model = SegmentationModel()
model.to(DEVICE)
optimizer = torch.optim.Adam(model.parameters(), lr=LR)
best_loss = np.Inf

for i in range(EPOCHS):
  # 取消训练阶段注释,正确调用训练函数
  train_loss = train_func(trainloader,model,optimizer) 
  valid_loss = eval_func(validloader,model)

  if valid_loss < best_loss:
    torch.save(model.state_dict(),"best-model.pt")
    print('SAVED')
    best_loss = valid_loss

  print(f"Epoch :  {i+1} Train Loss : {train_loss} Valid Loss : {valid_loss}")

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 23:00:57