Keras 3D U-Net多任务建模:Conv3D转平层分类报错求解
我尝试搭建一个可同时完成3D图像分割与分类的模型,用Keras函数式API构建3D U-Net,模型末尾设两个输出分支:一个做3D图像分割,另一个做癌症类型分类。单独跑分割任务时模型正常,但加了分类分支后报错。核心问题是如何把Conv3D输出转成Flatten层做Softmax分类,试过GlobalAveragePooling3D和Flatten()都没解决。
原模型代码
import segmentation_models_3D as sm from keras.models import Model from keras.layers import Input, Conv3D, MaxPooling3D, concatenate, Conv3DTranspose, BatchNormalization, Dropout, Lambda, Flatten, Dense , GlobalAveragePooling3D import keras.backend as K from tensorflow.keras.optimizers import Adam from keras.metrics import MeanIoU kernel_initializer = 'he_uniform' #Try others if you want def unet_model(IMG_HEIGHT, IMG_WIDTH, IMG_DEPTH, IMG_CHANNELS, num_classes): #Build the model inputs = Input((IMG_HEIGHT, IMG_WIDTH, IMG_DEPTH, IMG_CHANNELS), name="img") s = inputs #Contraction path c1 = Conv3D(16, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(s) c1 = Dropout(0.1)(c1) c1 = Conv3D(16, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c1) p1 = MaxPooling3D((2, 2, 2))(c1) c2 = Conv3D(32, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(p1) c2 = Dropout(0.1)(c2) c2 = Conv3D(32, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c2) p2 = MaxPooling3D((2, 2, 2))(c2) c3 = Conv3D(64, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(p2) c3 = Dropout(0.2)(c3) c3 = Conv3D(64, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c3) p3 = MaxPooling3D((2, 2, 2))(c3) c4 = Conv3D(128, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(p3) c4 = Dropout(0.2)(c4) c4 = Conv3D(128, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c4) p4 = MaxPooling3D(pool_size=(2, 2, 2))(c4) c5 = Conv3D(256, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(p4) c5 = Dropout(0.3)(c5) c5 = Conv3D(256, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c5) #Expansive path u6 = Conv3DTranspose(128, (2, 2, 2), strides=(2, 2, 2), padding='same')(c5) u6 = concatenate([u6, c4]) c6 = Conv3D(128, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(u6) c6 = Dropout(0.2)(c6) c6 = Conv3D(128, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c6) u7 = Conv3DTranspose(64, (2, 2, 2), strides=(2, 2, 2), padding='same')(c6) u7 = concatenate([u7, c3]) c7 = Conv3D(64, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(u7) c7 = Dropout(0.2)(c7) c7 = Conv3D(64, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c7) u8 = Conv3DTranspose(32, (2, 2, 2), strides=(2, 2, 2), padding='same')(c7) u8 = concatenate([u8, c2]) c8 = Conv3D(32, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(u8) c8 = Dropout(0.1)(c8) c8 = Conv3D(32, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c8) u9 = Conv3DTranspose(16, (2, 2, 2), strides=(2, 2, 2), padding='same')(c8) u9 = concatenate([u9, c1]) c9 = Conv3D(16, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(u9) c9 = Dropout(0.1)(c9) c9 = Conv3D(16, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c9) output_mask = Conv3D(num_classes, (1, 1, 1), activation='softmax', name ='mask')(c9) output_mask_1 = Conv3D(num_classes, (1, 1, 1), activation='softmax')(c9) output_label = GlobalAveragePooling3D()(output_mask_1) output_label = Dense(units=4, activation="relu")(output_label) output_label = Dense(units=3, activation="softmax",name="lab")(output_label) model = Model(inputs=inputs, outputs=[output_mask, output_label], name = "final_output") model.summary() return model
数据集生成器代码
def imageLoader(img_dir, img_list, mask_dir, mask_list, label_dir, label_list, batch_size): # some changes L = 5#len(img_list) #keras needs the generator infinite, so we will use while true while True: batch_start = 0 batch_end = batch_size # L-1 while batch_start < L: limit = min(batch_end, L) X = load_img(img_dir, img_list[batch_start:limit]) Y = load_img(mask_dir, mask_list[batch_start:limit]) # add classification label. yield ({"img":X}, {"mask":Y,"lab":np.array([[0,1,0],[0,1,0]]) }) # one hot encoding for 3 class classification. batch_start += batch_size batch_end += batch_size
原配置代码
import segmentation_models_3D as sm dice_loss = sm.losses.DiceLoss(class_weights=np.array([wt0, wt1, wt2, wt3])) focal_loss = sm.losses.CategoricalFocalLoss() total_loss = dice_loss + (1 * focal_loss) metrics = [ sm.metrics.IOUScore(threshold=0.5), 'accuracy'] LR = 0.0001 optim=tf.keras.optimizers.Adam(LR) model = unet_model(IMG_HEIGHT=128, IMG_WIDTH=128, IMG_DEPTH=128, IMG_CHANNELS=3, num_classes=4) model.compile(optimizer = optim, loss={"mask":total_loss, "lab": keras.losses.BinaryCrossentropy() } , metrics=metrics) print(model.summary())
报错信息
ValueError: in user code:
File "/usr/local/lib/python3.10/dist-packages/keras/engine/training.py", line 1021, in train_function *
return step_function(self, iterator)
File "/usr/local/lib/python3.10/dist-packages/segmentation_models_3D/metrics.py", line 62, in call *
**self.submodules
File "/usr/local/lib/python3.10/dist-packages/segmentation_models_3D/base/functional.py", line 93, in iou_score *
intersection = backend.sum(gt * pr, axis=axes)
File "/usr/local/lib/python3.10/dist-packages/keras/backend.py", line 2544, in sum
return tf.reduce_sum(x, axis, keepdims)
ValueError: Invalid reduction dimension 2 for input with 2 dimensions. for '{{node Sum_10}} = Sum[T=DT_FLOAT, Tidx=DT_INT32, keep_dims=false](mul_5, Sum_10/reduction_indices)' with input shapes: [?,3], [4] and with computed input tensors: input1 = <0 1 2 3>.
问题根源与解决方法
报错原因
你给所有输出分支统一指定了分割任务的指标(IOUScore),但这个指标是为3D分割的高维输出(形状[batch, H, W, D, num_classes])设计的,它会尝试对H、W、D对应的维度(轴1、2、3)求和计算IOU。而分类分支的输出是二维数组[batch, 3],没有这些维度,因此求和时触发维度不匹配的错误。
同时还有两个小问题:
- 3分类任务应该用
CategoricalCrossentropy损失,不是BinaryCrossentropy - 分类分支从分割输出提取特征冗余,直接从U-Net的瓶颈层(c5)提取更合理
修改后的代码
1. 修正模型定义(分类分支从瓶颈层取特征)
import segmentation_models_3D as sm from keras.models import Model from keras.layers import Input, Conv3D, MaxPooling3D, concatenate, Conv3DTranspose, Dropout, Dense, GlobalAveragePooling3D from tensorflow.keras.optimizers import Adam kernel_initializer = 'he_uniform' def unet_model(IMG_HEIGHT, IMG_WIDTH, IMG_DEPTH, IMG_CHANNELS, seg_num_classes, cls_num_classes): inputs = Input((IMG_HEIGHT, IMG_WIDTH, IMG_DEPTH, IMG_CHANNELS), name="img") s = inputs # 编码器部分(不变) c1 = Conv3D(16, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(s) c1 = Dropout(0.1)(c1) c1 = Conv3D(16, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c1) p1 = MaxPooling3D((2, 2, 2))(c1) c2 = Conv3D(32, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(p1) c2 = Dropout(0.1)(c2) c2 = Conv3D(32, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c2) p2 = MaxPooling3D((2, 2, 2))(c2) c3 = Conv3D(64, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(p2) c3 = Dropout(0.2)(c3) c3 = Conv3D(64, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c3) p3 = MaxPooling3D((2, 2, 2))(c3) c4 = Conv3D(128, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(p3) c4 = Dropout(0.2)(c4) c4 = Conv3D(128, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c4) p4 = MaxPooling3D(pool_size=(2, 2, 2))(c4) c5 = Conv3D(256, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(p4) c5 = Dropout(0.3)(c5) c5 = Conv3D(256, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c5) # 解码器部分(不变) u6 = Conv3DTranspose(128, (2, 2, 2), strides=(2, 2, 2), padding='same')(c5) u6 = concatenate([u6, c4]) c6 = Conv3D(128, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(u6) c6 = Dropout(0.2)(c6) c6 = Conv3D(128, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c6) u7 = Conv3DTranspose(64, (2, 2, 2), strides=(2, 2, 2), padding='same')(c6) u7 = concatenate([u7, c3]) c7 = Conv3D(64, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(u7) c7 = Dropout(0.2)(c7) c7 = Conv3D(64, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c7) u8 = Conv3DTranspose(32, (2, 2, 2), strides=(2, 2, 2), padding='same')(c7) u8 = concatenate([u8, c2]) c8 = Conv3D(32, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(u8) c8 = Dropout(0.1)(c8) c8 = Conv3D(32, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c8) u9 = Conv3DTranspose(16, (2, 2, 2), strides=(2, 2, 2), padding='same')(c8) u9 = concatenate([u9, c1]) c9 = Conv3D(16, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(u9) c9 = Dropout(0.1)(c9) c9 = Conv3D(16, (3, 3, 3), activation='relu', kernel_initializer=kernel_initializer, padding='same')(c9) # 分割输出分支 output_mask = Conv3D(seg_num_classes, (1, 1, 1), activation='softmax', name='mask')(c9) # 分类输出分支:从瓶颈层c5提取特征 cls_features = GlobalAveragePooling3D()(c5) cls_features = Dense(64, activation='relu')(cls_features) cls_features = Dropout(0.3)(cls_features) output_label = Dense(cls_num_classes, activation='softmax', name='lab')(cls_features) model = Model(inputs=inputs, outputs=[output_mask, output_label], name="final_output") model.summary() return model

