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

自定义Keras Sequence生成器触发NotImplementedError求助

解决Keras Sequence自定义生成器的NotImplementedError问题

错误原因

从报错日志可以明确看到,问题源于未实现__len__方法——继承tf.keras.utils.Sequence必须强制重写__init__、__len__、__getitem__三个核心方法;on_epoch_end是可选方法,但如果你的生成器需要在每个epoch结束后打乱数据,也必须实现它。

修复步骤

1. 实现__len__方法

该方法需要返回生成器的总批次数,计算逻辑为总样本数除以批次大小后向上取整(处理最后一批样本不足批次大小的情况)。

2. 实现on_epoch_end方法

如果初始化时设置了suffle=True,这个方法会在每个epoch结束后打乱样本顺序,保证训练随机性。注意要同时打乱图像和掩码的文件名,确保两者一一对应。

3. 修复__getitem__中的索引问题

原代码读取文件时的索引逻辑存在错位风险,应直接使用当前批次的文件名列表,而非通过index * self.batch_size + i计算索引。

修改后的完整代码

import math
import numpy as np
import os
import tensorflow as tf

class DataGenerator(tf.keras.utils.Sequence):
    def __init__(self, root_dir=r'../data/val_test', image_folder='img/', mask_folder='masks/', 
                 batch_size=4, image_size=288, nb_y_features=1, 
                 augmentation=None,
                 suffle=True):
        # 初始化图像和掩码文件名列表
        image_dir = os.path.join(root_dir, image_folder)
        mask_dir = os.path.join(root_dir, mask_folder)
        self.image_filenames = np.sort([os.path.join(image_dir, f) for f in os.listdir(image_dir)])
        self.mask_names = np.sort([os.path.join(mask_dir, f) for f in os.listdir(mask_dir)])

        self.batch_size = batch_size
        self.augmentation = augmentation
        self.image_size = image_size
        self.nb_y_features = nb_y_features
        self.suffle = suffle

        # 初始化时打乱数据(如果需要)
        if self.suffle:
            self.on_epoch_end()

    def __len__(self):
        # 返回总批次数,向上取整
        return math.ceil(len(self.image_filenames) / self.batch_size)

    def on_epoch_end(self):
        # 打乱图像和掩码的对应索引
        if self.suffle:
            indexes = np.arange(len(self.image_filenames))
            np.random.shuffle(indexes)
            self.image_filenames = self.image_filenames[indexes]
            self.mask_names = self.mask_names[indexes]

    def __getitem__(self, index):
        # 计算当前批次的样本索引范围
        start_idx = index * self.batch_size
        end_idx = min((index + 1) * self.batch_size, len(self.image_filenames))
        batch_image_filenames = self.image_filenames[start_idx:end_idx]
        batch_mask_names = self.mask_names[start_idx:end_idx]
        this_batch_size = len(batch_image_filenames)

        # 初始化批次数据数组
        X = np.empty((this_batch_size, self.image_size, self.image_size, 3), dtype=np.float32)
        y = np.empty((this_batch_size, self.image_size, self.image_size, self.nb_y_features), dtype=np.uint8)

        # 读取并处理每个样本
        for i in range(this_batch_size):
            img_path = batch_image_filenames[i]
            mask_path = batch_mask_names[i]
            X_sample, y_sample = self.read_image_mask(img_path, mask_path)

            if self.augmentation is not None:
                # 数据增强
                augmented = self.augmentation(self.image_size)(image=X_sample, mask=y_sample)
                image_augm = augmented['image']
                mask_augm = augmented['mask'].reshape(self.image_size, self.image_size, self.nb_y_features)
                # 归一化到0-1
                X[i, ...] = image_augm / 255.0
                y[i, ...] = mask_augm / 255.0
            else:
                # 验证集/测试集直接归一化
                X[i, ...] = X_sample / 255.0
                y[i, ...] = y_sample.reshape(self.image_size, self.image_size, self.nb_y_features) / 255.0

        return X, y

    # 补充read_image_mask方法(需自行实现文件读取逻辑)
    def read_image_mask(self, img_path, mask_path):
        # 示例实现:读取图像和掩码
        image = tf.keras.preprocessing.image.load_img(img_path, target_size=(self.image_size, self.image_size))
        image = tf.keras.preprocessing.image.img_to_array(image)
        mask = tf.keras.preprocessing.image.load_img(mask_path, target_size=(self.image_size, self.image_size), color_mode='grayscale')
        mask = tf.keras.preprocessing.image.img_to_array(mask)
        return image, mask

额外注意事项

  • 确保read_image_mask方法已正确实现,负责读取图像和掩码文件并返回符合尺寸要求的数组。
  • 若无需数据增强,augmentation参数传None即可,代码会自动走else分支处理。

内容的提问来源于stack exchange,提问作者Samin Zeini

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.12 22:25:41