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

如何在CustomDataGenerator中获取y_true并构建混淆矩阵?

问题分析与解决方案

错误ValueError: Found input variables with inconsistent numbers of samples: [0, 62]的核心原因是:你定义的self.y_true实例变量从未被实际填充,调用get_y_true()返回的是空列表,而y_pred有62个样本,两者数量不匹配。同时你的数据生成器还存在图像与标签配对错位的隐藏问题,以下是针对性修复:


1. 修复CustomDataGenerator代码

核心修改点:

  • 将读取到的标签存入实例的self.y_true(而非局部变量)
  • 修复图像与标签的打乱逻辑,避免配对错位
  • 处理最后一个batch样本数不足的情况
class CustomDataGenerator(Sequence):
    def __init__(self, image_folders, label_folders, dir, dim=(512,512),  batch_size=1,n_classes=7,n_channels=8,shuffle=True):
        self.image_folders = image_folders
        self.label_folders = label_folders
        self.dir = dir
        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 = []
        self.y_true = []
        self.on_epoch_end()

    def __len__(self):
        return int(np.ceil(len(self.image_paths) / self.batch_size))  

    def __getitem__(self, index):
        batch_image_paths = self.image_paths[index * self.batch_size: (index + 1) * self.batch_size]
        batch_label_paths = self.label_paths[index * self.batch_size: (index + 1) * self.batch_size]
        batch = zip(batch_image_paths, batch_label_paths)
        return self.get_data(batch)

    def on_epoch_end(self):
        self.image_paths = []
        self.label_paths = []
        self.y_true = []  # 每个epoch开始时清空标签列表
        
        # 收集图像路径
        for folder in self.image_folders:
            image_folder_path = os.path.join(self.dir, 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(self.dir, 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))
        
        # 修复:配对图像与标签后再打乱,避免错位
        if self.shuffle:
            paired = list(zip(self.image_paths, self.label_paths))
            np.random.shuffle(paired)
            self.image_paths, self.label_paths = zip(*paired)
            self.image_paths = list(self.image_paths)
            self.label_paths = list(self.label_paths)

    def get_data(self, batch):
        # 用len(batch)替代batch_size,处理最后一个batch样本数不足的情况
        X = np.empty((len(batch), *self.dim, self.n_channels))
        y = np.empty((len(batch), self.n_classes))

        for i, (image_path, label_path) in enumerate(batch):
            image = np.load(image_path)
            with open(label_path, 'r') as f:
                line = f.readline().strip()
                filepath, label = line.rsplit(' ', 1)
                label = int(label)
                self.y_true.append(label)  # 将标签存入实例变量
            label_one_hot = to_categorical(label, num_classes=self.n_classes)

            X[i,] = image
            y[i,] = label_one_hot

        return X, y
    
    def get_y_true(self):
        return self.y_true

2. 正确获取y_true并构建混淆矩阵

注意事项:

  • 验证集建议关闭shuffle,保证标签与预测结果顺序对应
  • 先调用model.predict()遍历验证集,此时生成器会自动填充self.y_true
train_datagen = CustomDataGenerator(image_folders, label_folders, train_dir, **params, shuffle=True)
# 验证集关闭shuffle,避免标签顺序混乱
val_datagen = CustomDataGenerator(image_folders, label_folders, valid_dir, **params, shuffle=False)

# 先执行预测,遍历验证集填充y_true
Y_pred = model.predict(val_datagen, verbose=1)
y_pred = np.argmax(Y_pred, axis=1) 

# 获取填充后的真实标签
y_true = val_datagen.get_y_true()

# 可先验证样本数量是否匹配
print(f"真实标签样本数:{len(y_true)}, 预测结果样本数:{len(y_pred)}")

# 构建混淆矩阵
sns.heatmap(confusion_matrix(y_true, y_pred), annot=True, fmt="d", cmap='Greens', ax=ax)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 16:35:00