timm使用Mixup时报Batch size should be even错误排查
问题原因
报错出现在训练轮次的最后一个批次(进度条显示333/334处),和你设置的batch_size=32是偶数没有关系:
- timm的Mixup/Cutmix实现需要将单个批次内的样本两两配对做混合,强制要求输入的每个batch的样本数为偶数
- 当训练集总样本数无法被batch size整除时,PyTorch的DataLoader默认会返回最后一个不足设定batch size的批次,这个批次的样本数刚好为奇数,就会触发断言报错。
解决方案
选任意一种即可:
方案1:训练时丢弃最后一个不完整批次(推荐)
调用trainer.train()时传入训练集DataLoader的配置,开启drop_last,直接丢弃最后一个凑不满batch size的批次,改动最小,对训练效果几乎无影响:trainer.train( per_device_batch_size=batch_size, train_dataset=train_dataset, eval_dataset=eval_dataset, num_epochs=num_epochs, create_scheduler_fn=trainer.create_scheduler, # 新增下面这行配置 train_dataloader_kwargs={"drop_last": True} )方案2:手动过滤奇数长度的批次
如果你不想丢弃最后一部分样本,可以修改calculate_train_batch_loss方法,遇到奇数长度的batch时主动去掉最后一个样本再传入Mixup:def calculate_train_batch_loss(self, batch): xb, yb = batch # 检查当前batch长度,奇数则截掉最后一个样本 if xb.size(0) % 2 != 0: xb = xb[:-1] yb = yb[:-1] mixup_xb, mixup_yb = self.mixup_fn(xb, yb) return super().calculate_train_batch_loss((mixup_xb, mixup_yb))方案3:调整配置避免不完整批次
可以通过调整batch size、对训练集做少量样本增删,让训练集总样本数可以被(单卡batch size * 分布式卡数)整除,从根源上避免出现不完整批次,不过这个方案灵活性差,一般不推荐。
内容的提问来源于stack exchange,提问作者Khawar Islam
相关产品推荐
相关产品推荐

