如何为U-Net添加额外跳跃连接路径实现双CT窗输入分割
双输入窗CT图像分割U-Net实现方案
你需要的双独立输入路径实现逻辑非常直接:给第二组不同窗宽窗位的CT图像单独搭建一条和原有编码路径(下采样部分)完全对称的特征提取分支,在每个下采样层级将两个分支输出的同尺度特征做通道拼接,再送入后续的解码上采样路径即可,完全匹配你标注的红色新增结构逻辑。
核心修改点
- 新增第二个输入层,对应第二组窗设置的CT图像,同样做归一化预处理
- 复刻原有的下采样编码结构作为第二分支,权重独立不共享,专门提取第二组CT图像的特征
- 在每个池化操作前的特征输出位置,把两个分支同尺度的特征沿通道维度拼接,作为下一层下采样、以及后续上采样跳连的输入
- 模型初始化时传入两个输入层组成的列表,训练时按照
[第一组CT图像数组, 第二组CT图像数组]的格式喂入数据即可,标签还是原来的单份分割掩码不用改
修改后可直接运行的代码
import tensorflow as tf IMG_WIDTH = 128 IMG_HEIGHT = 128 IMG_CHANNELS = 3 # 构建双输入模型 # 原第一组窗设置图像输入 input1 = tf.keras.layers.Input((IMG_HEIGHT, IMG_WIDTH, IMG_CHANNELS)) s1 = tf.keras.layers.Lambda(lambda x: x / 255)(input1) # 新增第二组窗设置图像输入 input2 = tf.keras.layers.Input((IMG_HEIGHT, IMG_WIDTH, IMG_CHANNELS)) s2 = tf.keras.layers.Lambda(lambda x: x / 255)(input2) # 下采样路径 第一分支(原分支) c1_1 = tf.keras.layers.Conv2D(16, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(s1) c1_1 = tf.keras.layers.Dropout(0.1)(c1_1) c1_1 = tf.keras.layers.Conv2D(16, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c1_1) p1_1 = tf.keras.layers.MaxPooling2D((2, 2))(c1_1) c2_1 = tf.keras.layers.Conv2D(32, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(p1_1) c2_1 = tf.keras.layers.Dropout(0.1)(c2_1) c2_1 = tf.keras.layers.Conv2D(32, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c2_1) p2_1 = tf.keras.layers.MaxPooling2D((2, 2))(c2_1) c3_1 = tf.keras.layers.Conv2D(64, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(p2_1) c3_1 = tf.keras.layers.Dropout(0.2)(c3_1) c3_1 = tf.keras.layers.Conv2D(64, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c3_1) p3_1 = tf.keras.layers.MaxPooling2D((2, 2))(c3_1) c4_1 = tf.keras.layers.Conv2D(128, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(p3_1) c4_1 = tf.keras.layers.Dropout(0.2)(c4_1) c4_1 = tf.keras.layers.Conv2D(128, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c4_1) p4_1 = tf.keras.layers.MaxPooling2D(pool_size=(2, 2))(c4_1) c5_1 = tf.keras.layers.Conv2D(256, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(p4_1) c5_1 = tf.keras.layers.Dropout(0.3)(c5_1) c5_1 = tf.keras.layers.Conv2D(256, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c5_1) # 下采样路径 第二分支(新增红色标注分支) c1_2 = tf.keras.layers.Conv2D(16, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(s2) c1_2 = tf.keras.layers.Dropout(0.1)(c1_2) c1_2 = tf.keras.layers.Conv2D(16, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c1_2) p1_2 = tf.keras.layers.MaxPooling2D((2, 2))(c1_2) c2_2 = tf.keras.layers.Conv2D(32, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(p1_2) c2_2 = tf.keras.layers.Dropout(0.1)(c2_2) c2_2 = tf.keras.layers.Conv2D(32, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c2_2) p2_2 = tf.keras.layers.MaxPooling2D((2, 2))(c2_2) c3_2 = tf.keras.layers.Conv2D(64, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(p2_2) c3_2 = tf.keras.layers.Dropout(0.2)(c3_2) c3_2 = tf.keras.layers.Conv2D(64, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c3_2) p3_2 = tf.keras.layers.MaxPooling2D((2, 2))(c3_2) c4_2 = tf.keras.layers.Conv2D(128, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(p3_2) c4_2 = tf.keras.layers.Dropout(0.2)(c4_2) c4_2 = tf.keras.layers.Conv2D(128, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c4_2) p4_2 = tf.keras.layers.MaxPooling2D(pool_size=(2, 2))(c4_2) c5_2 = tf.keras.layers.Conv2D(256, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(p4_2) c5_2 = tf.keras.layers.Dropout(0.3)(c5_2) c5_2 = tf.keras.layers.Conv2D(256, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c5_2) # 同层级特征拼接,作为瓶颈层和跳连输入 c1 = tf.keras.layers.concatenate([c1_1, c1_2], axis=3) c2 = tf.keras.layers.concatenate([c2_1, c2_2], axis=3) c3 = tf.keras.layers.concatenate([c3_1, c3_2], axis=3) c4 = tf.keras.layers.concatenate([c4_1, c4_2], axis=3) c5 = tf.keras.layers.concatenate([c5_1, c5_2], axis=3) # 上采样解码路径(和原结构逻辑一致,通道数随拼接特征自动适配) u6 = tf.keras.layers.Conv2DTranspose(128, (2, 2), strides=(2, 2), padding='same')(c5) u6 = tf.keras.layers.concatenate([u6, c4]) c6 = tf.keras.layers.Conv2D(128, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(u6) c6 = tf.keras.layers.Dropout(0.2)(c6) c6 = tf.keras.layers.Conv2D(128, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c6) u7 = tf.keras.layers.Conv2DTranspose(64, (2, 2), strides=(2, 2), padding='same')(c6) u7 = tf.keras.layers.concatenate([u7, c3]) c7 = tf.keras.layers.Conv2D(64, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(u7) c7 = tf.keras.layers.Dropout(0.2)(c7) c7 = tf.keras.layers.Conv2D(64, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c7) u8 = tf.keras.layers.Conv2DTranspose(32, (2, 2), strides=(2, 2), padding='same')(c7) u8 = tf.keras.layers.concatenate([u8, c2]) c8 = tf.keras.layers.Conv2D(32, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(u8) c8 = tf.keras.layers.Dropout(0.1)(c8) c8 = tf.keras.layers.Conv2D(32, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c8) u9 = tf.keras.layers.Conv2DTranspose(16, (2, 2), strides=(2, 2), padding='same')(c8) u9 = tf.keras.layers.concatenate([u9, c1], axis=3) c9 = tf.keras.layers.Conv2D(16, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(u9) c9 = tf.keras.layers.Dropout(0.1)(c9) c9 = tf.keras.layers.Conv2D(16, (3, 3), activation='relu', kernel_initializer='he_normal', padding='same')(c9) outputs = tf.keras.layers.Conv2D(1, (1, 1), activation='sigmoid')(c9) model = tf.keras.Model(inputs=[input1, input2], outputs=[outputs]) model.compile(optimizer='adam', loss='binary_crossentropy', metrics=['accuracy']) model.summary()
训练注意事项
- 喂入数据时需要保证两组图像是一一配对的,即同一解剖位置的两个不同窗设置图像组成样本对,顺序和输入层定义顺序对应
- 分割掩码不需要做任何修改,和单输入时完全一致
- 如果训练时发现显存占用过高,可以适当把两个分支的卷积核数量减半,或者降低输入分辨率,双分支结构本身不会改变模型训练的其他超参数设置逻辑
内容的提问来源于stack exchange,提问作者sumit
相关产品推荐
相关产品推荐

