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

多任务学习场景下,如何正确定义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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 08:45:48