多任务学习场景下,如何正确定义Keras的Train/Test生成器?
解决多任务学习中Train/Test生成器的定义问题
问题核心:flow_from_dataframe使用class_mode="raw"时,会返回包含所有目标列的单个数组,但你的模型需要两个独立的输出张量——一个对应分类任务的expr,另一个对应回归任务的valence+arousal,两者格式不匹配导致报错。以下是两种可行的解决方案:
方法一:封装现有生成器,拆分输出
基于你已有的flow_from_dataframe生成器,添加一层封装来拆分目标变量,匹配模型的双输出结构:
# 1. 创建基础生成器,获取图像和所有目标列 base_train_gen = img_gen.flow_from_dataframe( dataframe=train_dataset, x_col="file_loc", y_col=["expr", "valence", "arousal"], target_size=(96, 96), batch_size=203, class_mode="raw", shuffle=True ) base_test_gen = img_gen.flow_from_dataframe( dataframe=test_dataset_va, x_col="file_loc", y_col=["expr", "valence", "arousal"], target_size=(96, 96), batch_size=93, class_mode="raw", shuffle=False ) # 2. 自定义生成器函数,拆分目标为两个部分 def multi_task_generator(base_gen): for x, y in base_gen: # y的结构为 (batch_size, 3):第0列是expr,第1-2列是valence/arousal expr_y = y[:, 0] # 分类目标,形状(batch_size,) va_y = y[:, 1:] # 回归目标,形状(batch_size,2) yield x, [expr_y, va_y] # 3. 生成最终的训练/测试生成器 train_generator = multi_task_generator(base_train_gen) test_generator = multi_task_generator(base_test_gen)
使用时直接传入model.fit_generator即可,生成器返回的格式会完全匹配模型的双输出需求。
方法二:用tf.data.Dataset构建生成器(更推荐)
TensorFlow的Dataset API灵活性更高,适合多任务场景:
import tensorflow as tf def load_image(file_path): # 读取并预处理图像,和flow_from_dataframe保持一致 img = tf.io.read_file(file_path) img = tf.image.decode_jpeg(img, channels=3) img = tf.image.resize(img, (96, 96)) # 应用ResNet的预处理规则 img = tf.keras.applications.resnet50.preprocess_input(img) return img # 构建训练数据集 train_ds = tf.data.Dataset.from_tensor_slices( ( train_dataset["file_loc"].values, (train_dataset["expr"].values, train_dataset[["valence", "arousal"]].values) ) ) train_ds = train_ds.map( lambda x, y: (load_image(x), y), num_parallel_calls=tf.data.AUTOTUNE ).shuffle(len(train_dataset)).batch(203).prefetch(tf.data.AUTOTUNE) # 构建测试数据集 test_ds = tf.data.Dataset.from_tensor_slices( ( test_dataset_va["file_loc"].values, (test_dataset_va["expr"].values, test_dataset_va[["valence", "arousal"]].values) ) ) test_ds = test_ds.map( lambda x, y: (load_image(x), y), num_parallel_calls=tf.data.AUTOTUNE ).batch(93).prefetch(tf.data.AUTOTUNE)
使用时替换fit_generator为fit:
history = model.fit( train_ds, epochs=2, validation_data=test_ds, verbose=1 )
注意事项
- 若
img_gen包含数据增强(如翻转、旋转),方法一中的base_train_gen会自动应用这些操作,无需额外处理;方法二中可在load_image后添加增强逻辑,或把增强层加入模型。 - 确保
expr列是整数类型(匹配sparse_categorical_crossentropy需求),若不是需先转换:train_dataset["expr"] = train_dataset["expr"].astype(int)。
内容的提问来源于stack exchange,提问作者Alexandra Kapa
相关产品推荐
相关产品推荐

