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

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

关键修改说明

  1. 新增索引数组:在__init__中初始化self.indexes为路径数组的索引序列,用于后续打乱操作
  2. 修改on_epoch_end:仅打乱self.indexes,不再覆盖存储路径的self.image_paths和self.label_paths
  3. 调整__getitem__:通过打乱后的索引从原路径数组中获取对应的文件路径,确保传入np.load的是正确的字符串路径

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 21:37:04