tf.keras自定义Sequence数据生成器无法获取最后不完整批次求助
问题原因
__len__方法使用向下取整计算批次数量,直接舍去了不足一个batch的剩余样本__getitem__索引截取未做边界判断,最后一个批次如果样本不足会出现索引越界__data_generation方法固定按完整batch size初始化数组,无法适配最后一个小批次的实际尺寸
修复方案
1. 修改__len__方法,使用向上取整计算总批次
把原代码:
def __len__(self): 'Denotes the number of batches per epoch' num_batchs_per_epoch = int(np.floor(len(self.list_IDs) / self.batch_size)) return num_batchs_per_epoch
修改为:
def __len__(self): 'Denotes the number of batches per epoch' num_batchs_per_epoch = int(np.ceil(len(self.list_IDs) / self.batch_size)) return num_batchs_per_epoch
2. 修改__getitem__的索引边界判断,防止越界
把原代码:
end = (index+1)*self.batch_size
修改为:
end = min((index+1)*self.batch_size, len(self.indexes))
3. 修改__data_generation方法,按当前批次实际样本数初始化数组
把原代码:
X = np.empty((self.batch_size, 3, *self.dim, self.n_channels)) y1 = np.empty((self.batch_size), dtype=float) y2 = np.empty((self.batch_size), dtype=int)
修改为:
current_batch_size = len(list_IDs_temp) X = np.empty((current_batch_size, 3, *self.dim, self.n_channels)) y1 = np.empty((current_batch_size), dtype=float) y2 = np.empty((current_batch_size), dtype=int)
注意事项
如果你的模型中包含对批次大小有强约束的自定义算子,或是训练时小批次会显著影响BatchNormalization等层的效果,可以给生成器加一个drop_last参数,训练时设置为True丢弃最后不完整批次,验证/测试时设置为False保留全部样本。
内容的提问来源于stack exchange,提问作者user14909015
相关产品推荐
相关产品推荐

