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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.30 21:48:22