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,存在不必要的性能开销
修复方案
按以下步骤修改对应代码即可解决问题:
- 修正模型类,将损失函数初始化移到构造函数,统一变量名:
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
- 修正验证集增强,替换随机裁剪为固定尺寸缩放:
def get_val_augs(): return A.Compose([ A.Resize(width=IMG_SIZE, height=IMG_SIZE), ])
- 修正训练函数的变量名错误,移除函数内部错误缩进的调用代码:
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)
- 修正验证函数的变量名错误:
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)
- 修正数据集初始化、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
相关产品推荐
相关产品推荐

