TensorFlow中结合胸片与表格数据构建二分类模型的实现指导
多模态(胸片+表格数据)疾病二分类模型构建解决方案
关于Functional API的适用性
是,Functional API是这类多输入模型的最优方案。Sequential结构仅支持单输入单输出的线性层堆叠,无法处理图像、表格这种多分支的输入融合场景。而Functional API可以灵活定义多个输入分支、分别处理不同类型的数据,最后将特征融合,完美匹配你的需求。
原代码的错误分析与修正
核心错误点
- 缺失必要库导入:
LabelEncoder、train_test_split未导入 - 未定义
train_labels/test_labels:直接引用但未从数据集中提取 - 表格分支冗余的
Flatten:表格输入本身是一维向量,经过Dense后无需再Flatten - 输入数据不匹配:
ImageDataGenerator生成器与表格numpy数组的batch不同步,无法直接在fit中组合使用 - 标签处理错误:将标签转为字符串,但
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
相关产品推荐
相关产品推荐

