1. UNET图像语义分割基础认知

第一次接触UNET时,我盯着那个U型结构图看了整整半小时。这种对称的编码器-解码器设计,像极了小时候玩的拼图游戏——先把完整图片打散(下采样),再根据碎片特征重新拼合(上采样)。2015年诞生的UNET,最初是为医学影像分割设计的,但现在早已渗透到自动驾驶、卫星遥感等各个领域。

为什么UNET能成为语义分割的经典? 我总结出三个关键点:首先,跳跃连接(Skip Connection)就像施工队的对讲机,让底层和高层特征直接对话,避免了信息在传递过程中的损耗。其次,全卷积网络结构让它能处理任意尺寸的输入图像。最重要的是,在小样本数据上表现优异——这点在医疗领域尤其珍贵,毕竟标注一张CT切片可能需要医生数小时。

拿城市街景分割来说,我们需要区分道路、车辆、行人等不同对象。UNET会为每个像素打标签,就像用不同颜色的马克笔在照片上涂画。与普通分类网络不同,它的输出不是"这张图有车",而是精确到"这个像素属于车"。

2. 数据预处理实战技巧

2.1 数据集构建与标注

Cityscapes数据集是我的老熟人了,这个包含50个城市街景的数据集,标注文件就藏在_gtFine_labelIds.png里。但新手常会掉进坑里:训练集(train)和验证集(val)的目录结构必须严格一致。我建议用这个代码检查配对:

import os
from PIL import Image

train_images = sorted(glob.glob('leftImg8bit/train/*/*.png'))
train_masks = sorted(glob.glob('gtFine/train/*/*labelIds.png'))

# 检查文件名对应关系
for img_path, mask_path in zip(train_images[:5], train_masks[:5]):
    img_name = os.path.basename(img_path).replace('_leftImg8bit', '')
    mask_name = os.path.basename(mask_path)
    assert img_name == mask_name, f"文件名不匹配: {img_name} vs {mask_name}"

2.2 数据增强的陷阱

数据增强不是无脑操作!我曾因为随机翻转导致标注错位,模型精度直接腰斩。关键原则是:原图和标注必须同步变换。这里分享我的增强方案:

def augment(img, mask):
    if tf.random.uniform(()) > 0.5:
        img = tf.image.flip_left_right(img)
        mask = tf.image.flip_left_right(mask)
    
    if tf.random.uniform(()) > 0.5:
        img = tf.image.flip_up_down(img)
        mask = tf.image.flip_up_down(mask)
    
    # 随机亮度调整(仅对原图)
    img = tf.image.random_brightness(img, 0.2)
    return img, mask

特别注意:色彩变换只能用于原图,标注图必须保持像素值不变。曾经有学员把标注图也做了亮度调整,导致类别标签全部错乱,训练完全失败。

3. 模型构建核心细节

3.1 下采样模块设计

UNET的编码器部分像漏斗,逐步提取特征。但实现时有几个魔鬼细节:

  • 每个下采样块建议采用"卷积+BN+ReLU"的组合
  • 池化层建议用MaxPooling2D,核尺寸(2,2)
  • 通道数建议按64-128-256-512翻倍增长
def downsample_block(filters, size, apply_batchnorm=True):
    initializer = tf.keras.initializers.GlorotNormal()
    
    block = tf.keras.Sequential()
    block.add(tf.keras.layers.Conv2D(filters, size, 
                                    strides=1, padding='same',
                                    kernel_initializer=initializer))
    if apply_batchnorm:
        block.add(tf.keras.layers.BatchNormalization())
    block.add(tf.keras.layers.ReLU())
    
    return block

3.2 跳跃连接的秘密

UNET最精妙的就是跳跃连接。但新手常犯的错误是直接相加(add)——这会导致特征信息丢失。正确做法是通道维度拼接(concatenate):

def upsample_block(filters, size, apply_dropout=False):
    block = tf.keras.Sequential()
    block.add(tf.keras.layers.Conv2DTranspose(filters, size, 
                                             strides=2, padding='same'))
    if apply_dropout:
        block.add(tf.keras.layers.Dropout(0.5))
    return block

# 上采样时与跳跃连接合并
x = upsample_block(256, 3)(x)
x = tf.keras.layers.concatenate([x, skip_connection])

4. 训练优化策略

4.1 损失函数选择

语义分割不能用普通的交叉熵!我对比过三种损失函数:

  1. 加权交叉熵:对少数类别(如行人)增加权重
  2. Dice Loss:特别适合类别不平衡场景
  3. 复合损失:我的首选方案
def dice_coeff(y_true, y_pred, smooth=1):
    intersection = tf.reduce_sum(y_true * y_pred)
    return (2. * intersection + smooth) / (tf.reduce_sum(y_true) + tf.reduce_sum(y_pred) + smooth)

def dice_loss(y_true, y_pred):
    return 1 - dice_coeff(y_true, y_pred)

def total_loss(y_true, y_pred):
    return 0.7*dice_loss(y_true, y_pred) + 0.3*tf.keras.losses.binary_crossentropy(y_true, y_pred)

4.2 学习率调度

我用余弦退火+热重启的组合,比固定学习率提升约3% mIoU:

lr_schedule = tf.keras.optimizers.schedules.CosineDecayRestarts(
    initial_learning_rate=1e-3,
    first_decay_steps=1000,
    t_mul=2.0,
    m_mul=0.9
)

5. 模型部署实战

5.1 TensorRT加速

当需要实时推理时(如自动驾驶),我这样优化UNET:

# 转换模型为TensorRT
conversion_params = trt.TrtConversionParams(
    precision_mode=trt.TrtPrecisionMode.FP16,
    max_workspace_size_bytes=1 << 25
)

converter = trt.TrtGraphConverterV2(
    input_saved_model_dir='saved_model',
    conversion_params=conversion_params
)
converter.convert()
converter.save('trt_model')

5.2 移动端部署技巧

在安卓端部署时,发现原模型太大(120MB)。通过以下操作压缩到18MB:

  1. 量化训练(Quantization Aware Training)
  2. 通道剪枝(移除10%的冗余通道)
  3. TFLite转换时启用硬件加速
converter = tf.lite.TFLiteConverter.from_saved_model('saved_model')
converter.optimizations = [tf.lite.Optimize.DEFAULT]
converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS_INT8]
tflite_model = converter.convert()

6. 常见问题排坑指南

问题1:训练时loss震荡严重

  • 检查数据增强是否正确同步
  • 尝试减小初始学习率(如从1e-3降到5e-4)
  • 添加梯度裁剪:optimizer = tf.keras.optimizers.Adam(clipvalue=0.5)

问题2:预测结果有"雪花点"

  • 在最后一层卷积后添加CRF后处理
  • 尝试Dice系数阈值调优(通常0.3-0.5效果最佳)

问题3:显存不足

  • 降低batch size(可小至2)
  • 使用混合精度训练:
    policy = tf.keras.mixed_precision.Policy('mixed_float16')
    tf.keras.mixed_precision.set_global_policy(policy)
    

记得第一次成功分割出整条道路时,那种成就感至今难忘。UNET就像乐高积木,看似简单却能搭建出强大应用。现在每次看到自动驾驶车辆识别出人行道,都会想起那些调试模型的深夜。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐