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

tf.keras自定义Sequence数据生成器无法获取最后不完整批次求助

问题原因

  1. __len__方法使用向下取整计算批次数量,直接舍去了不足一个batch的剩余样本
  2. __getitem__索引截取未做边界判断,最后一个批次如果样本不足会出现索引越界
  3. __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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.02 20:06:05