自定义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
相关产品推荐
相关产品推荐

