Keras自定义数据生成器加载.npy文件报TypeError问题排查
问题解决:Keras自定义数据生成器np.load报错TypeError
错误根源
你在on_epoch_end方法中错误地将原本存储文件路径的self.image_paths和self.label_paths替换成了索引数组(np.arange(len(self.image_paths))),导致后续调用np.load(image_path)时,传入的是整数索引而非字符串路径,触发了类型错误。
修正后的代码
正确的做法是单独维护一个索引数组用于打乱顺序,而非覆盖原有的路径数组。修改后的CustomDataGenerator类如下:
class CustomDataGenerator(Sequence): def __init__(self, image_folders, label_folders,dim=(512,512), batch_size=1,n_classes=7,n_channels=18,shuffle=True): self.image_folders = image_folders self.label_folders = label_folders self.dim = dim self.batch_size = batch_size self.n_classes = n_classes self.n_channels = n_channels self.shuffle = shuffle self.image_paths = [] self.label_paths = [] for folder in self.image_folders: image_folder_path = os.path.join('data/syu/npy', folder) image_files = os.listdir(image_folder_path) for file_name in image_files: self.image_paths.append(os.path.join(image_folder_path, file_name)) for folder in self.label_folders: label_folder_path = os.path.join('data/syu/npy', folder) label_files = os.listdir(label_folder_path) for file_name in label_files: self.label_paths.append(os.path.join(label_folder_path, file_name)) # 初始化索引数组 self.indexes = np.arange(len(self.image_paths)) self.on_epoch_end() def __len__(self): return int(np.ceil(len(self.image_paths) / self.batch_size)) def __getitem__(self, index): # 通过打乱后的索引获取批次对应的路径 batch_indexes = self.indexes[index * self.batch_size: (index + 1) * self.batch_size] batch_image_paths = [self.image_paths[i] for i in batch_indexes] batch_label_paths = [self.label_paths[i] for i in batch_indexes] batch = zip(batch_image_paths, batch_label_paths) return self.get_data(batch) def on_epoch_end(self): # 仅打乱索引数组,不修改路径数组 if self.shuffle == True: np.random.shuffle(self.indexes) def get_data(self, batch): X = np.empty((self.batch_size, *self.dim, self.n_channels)) y = np.empty((self.batch_size, *self.dim, self.n_classes)) for i, (image_path, label_path) in enumerate(batch): image = np.load(image_path) label = np.load(label_path) X[i,] = image y[i,] = label return X, y
关键修改说明
- 新增索引数组:在
__init__中初始化self.indexes为路径数组的索引序列,用于后续打乱操作 - 修改
on_epoch_end:仅打乱self.indexes,不再覆盖存储路径的self.image_paths和self.label_paths - 调整
__getitem__:通过打乱后的索引从原路径数组中获取对应的文件路径,确保传入np.load的是正确的字符串路径
内容的提问来源于stack exchange,提问作者Syuuuu
相关产品推荐
相关产品推荐

