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

新冠胸片多分类模型切换损失函数后出现Graph Execution Error求助

新冠胸片多分类任务Graph Execution Error排查与修复

我正在参与新冠胸片多分类挑战赛,需将数据分为NOFINDING、COVID19、THORAXDISEASE三类,采用多分类方式更合理。但将损失函数从binary_crossentropy改为categorical_crossentropy后,调用model_pretrained.fit时持续出现Graph Execution Error。已设置IMG_SIZE=224,怀疑图像尺寸问题但未找到根源。

相关代码片段

数据加载部分

train_path = '/content/drive/MyDrive/covid19/csc532-2/DLAI3_Phase3/DLAI3_Phase3'

train_COVID19_1 = glob.glob(train_path+"/COVID-19/*.png")
train_NOFINDING_1 = glob.glob(train_path+"/NOFINDING/*.png")
train_THORAXDISEASE_1 = glob.glob(train_path+"/THORAXDISEASE/*.png")

train_COVID19_2 = glob.glob(train_path+"/COVID-19/*.jpg")
train_NOFINDING_2 = glob.glob(train_path+"/NOFINDING/*.jpg")
train_THORAXDISEASE_2 = glob.glob(train_path+"/THORAXDISEASE/*.jpg")

train_COVID19_3 = glob.glob(train_path+"/COVID-19/*.jpeg")
train_NOFINDING_3 = glob.glob(train_path+"/NOFINDING/*.jpeg")
train_THORAXDISEASE_3 = glob.glob(train_path+"/THORAXDISEASE/*.jpeg")
train_list = [x for x in train_COVID19_1]
train_list.extend([x for x in train_COVID19_2])
train_list.extend([x for x in train_COVID19_3])
train_list.extend([x for x in train_NOFINDING_1])
train_list.extend([x for x in train_NOFINDING_2])
train_list.extend([x for x in train_NOFINDING_3])
train_list.extend([x for x in train_THORAXDISEASE_1])
train_list.extend([x for x in train_THORAXDISEASE_2])
train_list.extend([x for x in train_THORAXDISEASE_3])

df_train = pd.DataFrame(np.concatenate([['COVID19']*(len(train_COVID19_1)+len(train_COVID19_2)+len(train_COVID19_3)), 
                                        ['NOFINDING']*(len(train_NOFINDING_1)+len(train_NOFINDING_2)+len(train_NOFINDING_3)),
                                        ['THORAXDISEASE']*(len(train_THORAXDISEASE_1)+len(train_THORAXDISEASE_2)+len(train_THORAXDISEASE_3))]), 
                        columns = ['class'])
df_train['image'] = [x for x in train_list]

数据生成器与模型编译

train_datagen = ImageDataGenerator(rescale=1/255.,
                                  zoom_range = 0.1,
                                  width_shift_range = 0.1,
                                  height_shift_range = 0.1)

val_datagen = ImageDataGenerator(rescale=1/255.)

ds_train = train_datagen.flow_from_dataframe(train_df,
                                             x_col = 'image',
                                             y_col = 'class',
                                             target_size = (IMG_SIZE, IMG_SIZE),
                                             class_mode = 'categorical',
                                             batch_size = BATCH,
                                             seed = SEED)

ds_val = val_datagen.flow_from_dataframe(test_df,
                                            x_col = 'image',
                                            y_col = 'class',
                                            target_size = (IMG_SIZE, IMG_SIZE),
                                            class_mode = 'categorical',
                                            batch_size = BATCH,
                                            seed = SEED)

ds_test = val_datagen.flow_from_dataframe(df_validate,
                                            x_col = 'image',
                                            y_col = 'class',
                                            target_size = (IMG_SIZE, IMG_SIZE),
                                            class_mode = 'categorical',
                                            batch_size = 1,
                                            shuffle = False)
keras.backend.clear_session()

model = get_model()
model.compile(loss='binary_crossentropy'
              , optimizer = keras.optimizers.Adam(learning_rate=3e-5), metrics='binary_accuracy')
 
model.summary()

触发错误的训练代码

history = model_pretrained.fit(ds_train,
          batch_size = BATCH, epochs = 30,
          validation_data=ds_val,
          callbacks=[early_stopping, plateau],
          steps_per_epoch=(len(train_df)/BATCH),
          validation_steps=(len(test_df)/BATCH));

核心问题排查与修复

出现该错误的根源基本不是图像尺寸,而是模型输出层与多分类任务不匹配或指标/数据格式不兼容,以下是具体修复步骤:

1. 修正模型输出层

使用categorical_crossentropy时,模型最后一层必须满足:

  • 激活函数为softmax
  • 神经元数量等于类别数(此处为3)

如果get_model()返回的模型是二分类设计,手动替换输出层:

# 移除原输出层并添加多分类输出层
model.layers.pop()
model.add(Dense(3, activation='softmax'))

2. 更换评估指标

原代码中的binary_accuracy是二分类指标,多分类任务需改用categorical_accuracy或通用accuracy:

model.compile(loss='categorical_crossentropy',
              optimizer=keras.optimizers.Adam(learning_rate=3e-5),
              metrics=['categorical_accuracy'])

3. 修复测试集数据格式

测试集df_validate的class列全为Validate,但设置了class_mode='categorical'会导致标签维度不匹配。无标签测试集需调整参数:

ds_test = val_datagen.flow_from_dataframe(df_validate,
                                          x_col='image',
                                          y_col=None,
                                          target_size=(IMG_SIZE, IMG_SIZE),
                                          class_mode=None,
                                          batch_size=1,
                                          shuffle=False)

4. 确认数据生成器类别一致性

检查训练/验证集的类别映射是否正确包含3类:

print(ds_train.class_indices)
# 正常输出应为 {'COVID19':0, 'NOFINDING':1, 'THORAXDISEASE':2}

5. 处理单通道灰度图

若胸片为单通道图像,需转换为预训练模型要求的3通道RGB:

def grayscale_to_rgb(img):
    return np.repeat(img, 3, axis=-1)

train_datagen = ImageDataGenerator(rescale=1/255.,
                                  zoom_range=0.1,
                                  width_shift_range=0.1,
                                  height_shift_range=0.1,
                                  preprocessing_function=grayscale_to_rgb)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 00:07:20