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

如何为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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.01 03:45:34