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

TensorFlow中结合胸片与表格数据构建二分类模型的实现指导

多模态(胸片+表格数据)疾病二分类模型构建解决方案

关于Functional API的适用性

是,Functional API是这类多输入模型的最优方案。Sequential结构仅支持单输入单输出的线性层堆叠,无法处理图像、表格这种多分支的输入融合场景。而Functional API可以灵活定义多个输入分支、分别处理不同类型的数据,最后将特征融合,完美匹配你的需求。

原代码的错误分析与修正

核心错误点

  1. 缺失必要库导入:LabelEncoder、train_test_split未导入
  2. 未定义train_labels/test_labels:直接引用但未从数据集中提取
  3. 表格分支冗余的Flatten:表格输入本身是一维向量,经过Dense后无需再Flatten
  4. 输入数据不匹配:ImageDataGenerator生成器与表格numpy数组的batch不同步,无法直接在fit中组合使用
  5. 标签处理错误:将标签转为字符串,但class_mode='binary'要求数值型标签

修正后的完整代码

import tensorflow as tf
from tensorflow.keras.layers import Input, Conv2D, MaxPooling2D, Flatten, Dense, concatenate
from tensorflow.keras.models import Model
# 补充缺失的库导入
from sklearn.preprocessing import LabelEncoder
from sklearn.model_selection import train_test_split
import numpy as np

# ---------------------- 数据预处理修正 ----------------------
# 标签保持数值型,无需转字符串
label_encoder = LabelEncoder()
df['disease'] = label_encoder.fit_transform(df['disease']).astype(np.float32)

# 拆分数据集
train_data, test_data = train_test_split(df, test_size=0.2, random_state=42, stratify=df['disease'])

# 图像生成器(注意class_mode='binary'要求标签是0/1数值)
image_datagen = ImageDataGenerator(rescale=1/255)

train_image_generator = image_datagen.flow_from_dataframe(
    train_data,
    x_col='filename',
    y_col='disease',
    target_size=(224, 224),
    batch_size=32,
    class_mode='raw',  # 使用raw直接返回原始数值标签,避免字符串转换问题
    dtype='float32',
    shuffle=False  # 关闭shuffle,保证和表格数据顺序一致
)

test_image_generator = image_datagen.flow_from_dataframe(
    test_data,
    x_col='filename',
    y_col='disease',
    target_size=(224, 224),
    batch_size=32,
    class_mode='raw',
    dtype='float32',
    shuffle=False
)

# 处理表格特征
train_data['Sex'] = train_data['Sex'].map({'F': 0, 'M': 1})
test_data['Sex'] = test_data['Sex'].map({'F': 0, 'M': 1})

train_tabular_features = train_data[['age', 'Sex', 'Height', 'Weight']].values.astype(np.float32)
test_tabular_features = test_data[['age', 'Sex', 'Height', 'Weight']].values.astype(np.float32)

# ---------------------- 自定义生成器:同步图像与表格数据 ----------------------
def combined_generator(image_gen, tabular_data):
    while True:
        # 获取图像batch和对应标签
        img_batch, label_batch = next(image_gen)
        # 获取当前batch的索引(因为image_gen按顺序输出,shuffle=False)
        batch_start = (image_gen.batch_index - 1) * image_gen.batch_size
        batch_end = batch_start + len(img_batch)
        # 提取对应索引的表格数据
        tab_batch = tabular_data[batch_start:batch_end]
        yield {'image_input': img_batch, 'tabular_input': tab_batch}, label_batch

train_combined_gen = combined_generator(train_image_generator, train_tabular_features)
test_combined_gen = combined_generator(test_image_generator, test_tabular_features)

# ---------------------- 模型结构修正 ----------------------
# 图像分支
img_input = Input(shape=(224, 224, 3), name='image_input')
x1 = Conv2D(16, 3, padding='same', activation='relu')(img_input)
x1 = MaxPooling2D()(x1)
x1 = Conv2D(32, 3, padding='same', activation='relu')(x1)
x1 = MaxPooling2D()(x1)
x1 = Flatten()(x1)

# 表格分支:移除冗余的Flatten
tabular_input = Input(shape=(4,), name='tabular_input')
x2 = Dense(16, activation='relu')(tabular_input)
x2 = Dense(32, activation='relu')(x2)

# 特征融合
concatenated = concatenate([x1, x2])

# 分类头
x = Dense(128, activation='relu')(concatenated)
output_layer = Dense(1, activation='sigmoid', name='output')(x)

model = Model(inputs=[img_input, tabular_input], outputs=output_layer)

# 编译模型
model.compile(
    optimizer='adam',
    loss='binary_crossentropy',
    metrics=['accuracy']
)

# ---------------------- 模型训练 ----------------------
# 计算steps_per_epoch和validation_steps
train_steps = len(train_image_generator)
val_steps = len(test_image_generator)

history = model.fit(
    train_combined_gen,
    steps_per_epoch=train_steps,
    epochs=20,
    validation_data=test_combined_gen,
    validation_steps=val_steps
)

关键修改说明

  • 库导入补充:添加了LabelEncoder和train_test_split的导入
  • 标签处理优化:使用class_mode='raw'直接返回数值标签,避免字符串转换导致的类型错误
  • 自定义生成器:解决图像生成器与表格数据的batch同步问题,保证每个图像batch对应正确的表格特征
  • 模型结构精简:移除表格分支中冗余的Flatten层,因为Dense层输出已经是一维特征
  • 训练流程调整:使用自定义生成器作为输入,设置正确的steps_per_epoch和validation_steps

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 06:11:15