tf.keras.utils.Sequence迭代时忽略小于批大小的最后一批数据
解决tf.keras.utils.Sequence迭代时忽略最后一批不足batch_size数据的问题
问题根源在你的TEST_DATA_GENERATOR类的__len__方法:
def __len__(self): return len(self.indices) // self.batch_size
这里用了整数除法,对于10个样本、batch_size=4的情况,10//4=2,迭代器只会循环2次,调用__getitem__(0)和__getitem__(1),直接跳过了第3批的2个样本。
解决方案:修改__len__方法为向上取整计算
把__len__改成向上取整的逻辑,确保所有样本都能被覆盖。可以用两种方式实现:
方式1:整数运算实现向上取整
def __len__(self): return (len(self.indices) + self.batch_size - 1) // self.batch_size
原理是通过加上batch_size-1,让余数部分触发进位,比如10+4-1=13,13//4=3,刚好得到正确的批次数量。
方式2:用math.ceil(需要导入math模块)
import math def __len__(self): return math.ceil(len(self.indices) / self.batch_size)
修改后的完整代码
import numpy as np import tensorflow as tf class TEST_DATA_GENERATOR(tf.keras.utils.Sequence): def __init__( self, ): self.samples = [1,2,3,4,5,6,7,8,9,10] # 原代码中model_config需提前定义,此处保留结构 # self.model_name = model_config["name"] self.batch_size = 4 self.shuffle = False self.indices = range(0, len(self.samples)) assert self.batch_size <= len(self.indices), "batch size must be smaller than the number of samples" self.on_epoch_end() # shuffle def __len__(self): # 修改为向上取整逻辑 return (len(self.indices) + self.batch_size - 1) // self.batch_size def __getitem__(self, index): index = self.index[index * self.batch_size:(index + 1) * self.batch_size] batch = [self.indices[k] for k in index] X, y = self.__get_data(batch) return X, y def on_epoch_end(self): self.index = np.arange(len(self.indices)) if self.shuffle == True: np.random.shuffle(self.index) def __get_data(self, batch): X = [] y = [] for i in range(len(batch)): y.append("classlabel") for batch_idx, sample_idx in enumerate(batch): X.append(self.samples[sample_idx]) X = np.asarray(X) y = np.asarray(y) return X, y testgen = TEST_DATA_GENERATOR() # 单独调用验证 x,y = testgen.__getitem__(0) print(x.shape) x,y = testgen.__getitem__(1) print(x.shape) x,y = testgen.__getitem__(2) print(x.shape) print("----") # 迭代验证 for x,y in testgen.__iter__(): print(x.shape)
修改后的输出
(4,) (4,) (2,) ---- (4,) (4,) (2,)
此时迭代器会返回所有批次,包括最后一批不足batch_size的样本。
内容的提问来源于stack exchange,提问作者Random4Logic
相关产品推荐
相关产品推荐

