如何构建适用于Keras的X_train与Y_train(多标签图像分类)
解决Keras多标签图像标注问题的具体实现方案
针对你的多独立二分类(多标签)图像标注场景,结合keras.preprocessing.image.ImageDataGenerator,我整理了一套完整的实现流程,分两种常用方案讲解:
核心前提说明
先明确多标签分类和普通多分类的关键区别,这是实现的核心:
- 输出层激活函数必须用sigmoid(每个标签独立输出0-1概率)
- 损失函数选择binary_crossentropy(每个标签单独做二分类损失计算)
- 稀疏矩阵格式的
Y_train需要转为密集数组(Keras生成器通常需要这种格式)
方案一:使用flow_from_dataframe(推荐,代码更简洁)
这种方法通过DataFrame关联图片路径和标签,适合已有标签矩阵的场景,实现起来最省心。
步骤1:准备数据映射
- 获取所有图片的路径列表,务必保证顺序和
Y_train的行顺序完全一致:
import os import glob import pandas as pd from scipy.sparse import csr_matrix # 假设图片存储在"images/"文件夹下 image_paths = sorted(glob.glob("images/*.jpg")) # 排序保证顺序稳定
- 将稀疏矩阵
Y_train转为密集数组,并构建DataFrame:
# 稀疏矩阵转密集数组 Y_dense = Y_train.toarray() # 生成标签名称(可替换为你的真实标签名) label_names = [f"label_{i}" for i in range(Y_dense.shape[1])] # 构建DataFrame:一列存图片文件名,其余列存对应标签的0/1值 df = pd.DataFrame(Y_dense, columns=label_names) df["filename"] = [os.path.basename(path) for path in image_paths]
步骤2:配置ImageDataGenerator
包含数据归一化、可选的增强操作,以及训练/验证集拆分:
from keras.preprocessing.image import ImageDataGenerator datagen = ImageDataGenerator( rescale=1./255, # 将像素值归一化到0-1区间 horizontal_flip=True, # 随机水平翻转,增强数据多样性 validation_split=0.2 # 拆分20%数据作为验证集 )
步骤3:创建训练/验证生成器
# 训练集生成器 train_generator = datagen.flow_from_dataframe( dataframe=df, directory="images/", # 图片所在的根文件夹 x_col="filename", # DataFrame中存储图片文件名的列名 y_col=label_names, # 所有标签列的名称列表 target_size=(224, 224), # 统一图片尺寸为224x224(可根据模型调整) batch_size=32, class_mode="raw", # 多标签分类必须用"raw",表示返回原始数值数组 subset="training" ) # 验证集生成器 val_generator = datagen.flow_from_dataframe( dataframe=df, directory="images/", x_col="filename", y_col=label_names, target_size=(224, 224), batch_size=32, class_mode="raw", subset="validation" )
步骤4:构建并训练模型
from keras.models import Sequential from keras.layers import Dense from keras.applications import ResNet50 # 用预训练模型作为特征提取器,也可以自定义卷积层 model = Sequential([ ResNet50(include_top=False, input_shape=(224,224,3), pooling="avg"), Dense(len(label_names), activation="sigmoid") # 输出层对应标签数量,sigmoid激活 ]) model.compile( optimizer="adam", loss="binary_crossentropy", # 多标签分类专用损失函数 metrics=["accuracy"] ) # 开始训练 model.fit( train_generator, validation_data=val_generator, epochs=10 )
方案二:自定义Sequence生成器(更灵活)
如果你的数据有特殊预处理需求,或者不想用DataFrame,可以继承keras.utils.Sequence写自定义生成器,这种方式可控性更强。
步骤1:实现自定义生成器类
from keras.utils import Sequence from keras.preprocessing.image import load_img, img_to_array import numpy as np from math import ceil class MultiLabelImageGenerator(Sequence): def __init__(self, image_paths, y_sparse, batch_size, target_size, datagen): self.image_paths = image_paths self.y_sparse = y_sparse self.batch_size = batch_size self.target_size = target_size self.datagen = datagen # 返回总批次数量 def __len__(self): return ceil(len(self.image_paths) / self.batch_size) # 生成单个批次的数据 def __getitem__(self, idx): # 获取当前批次的图片路径和标签 batch_start = idx * self.batch_size batch_end = min((idx+1)*self.batch_size, len(self.image_paths)) batch_paths = self.image_paths[batch_start:batch_end] batch_y = self.y_sparse[batch_start:batch_end].toarray() # 加载并预处理图片 batch_x = [] for path in batch_paths: # 加载图片并调整尺寸 img = load_img(path, target_size=self.target_size) # 转为数组 img_arr = img_to_array(img) # 应用数据增强(如翻转、旋转等) img_arr = self.datagen.random_transform(img_arr) # 标准化(如归一化) img_arr = self.datagen.standardize(img_arr) batch_x.append(img_arr) return np.array(batch_x), batch_y
步骤2:初始化生成器并训练
# 配置ImageDataGenerator datagen = ImageDataGenerator(rescale=1./255, horizontal_flip=True) # 初始化生成器 train_generator = MultiLabelImageGenerator( image_paths=image_paths, y_sparse=Y_train, batch_size=32, target_size=(224,224), datagen=datagen ) # 模型构建和训练和方案一完全一致 model = Sequential([ ResNet50(include_top=False, input_shape=(224,224,3), pooling="avg"), Dense(Y_train.shape[1], activation="sigmoid") ]) model.compile(optimizer="adam", loss="binary_crossentropy", metrics=["accuracy"]) model.fit(train_generator, epochs=10)
关键注意事项
- 顺序一致性:图片路径的顺序必须和
Y_train的行顺序严格对应,否则标签会错位,建议用sorted()对文件名排序。 - 图片尺寸:必须指定
target_size,因为Keras模型需要固定尺寸的输入,可根据你选择的模型调整(比如VGG16用224x224,EfficientNet用224/240/384等)。 - 标签处理:稀疏矩阵必须转为密集数组,因为Keras生成器无法直接处理稀疏矩阵格式的标签。
内容的提问来源于stack exchange,提问作者Michael
相关产品推荐
相关产品推荐

