恒美微站
首页
关于我们
建站服务
主题模板
案例展示
资讯中心
联系我们
Unet与SAM结合的车道线分割实战:从提示框热图到训练调参
首页
资讯中心
/
Unet与SAM结合的车道线分割实战:从提示框热图到训练调参
Unet与SAM结合的车道线分割实战:从提示框热图到训练调参
发布时间:2026/9/10 16:21:08
简介该资源面向自动驾驶感知与语义分割方向的算法工程师、研究人员及高校学生提供一套基于Unet结合SAM提示框机制的端到端道路线分割方案。方案创新性地将SAM的提示引导能力引入Unet编码器针对道路线弱语义、细长形态目标进行专门优化适合真实驾驶场景下的车道线识别与可行驶区域划分。压缩包共2000个文件其中以png与jpg格式的标注图像为主1991个覆盖多时段、多光照条件下道路场景另有5个Python源码文件包含模型训练、推理与评估脚本以及3个txt配置文件和1个readme说明文档。资源包整体约530MB目录结构规整便于按数据类型与代码模块快速检索。作者以200个epoch完成训练验证集Dice系数达到0.89左右具备较强分割精度读者可直接使用预训练权重进行迁移学习或在自有数据上微调减少重复训练成本。已有639人学习下载适合需要快速搭建道路线分割基线或深入理解Unet与SAM融合策略的研究者。1. 为什么是 UnetSAM道路线分割不能只靠一种网络完成自动驾驶场景里的“道路线分割”和常规语义分割有个明显差别车道线是窄、长、连续性极强的结构往往只占图像面积的 0.5%~2%类间极度不平衡。用纯 Unet 训练模型容易把精力全放在沥青路面和背景上车道线 mIoU 常常卡在 40% 以下换成 SAM 直接推理它又把路沿、路缘石、斑马线全当成“前景物体”分割出来拿不到语义标签。把两者接在一起让 SAM 的提示框先给出几何先验再让 Unet 做语义决策是当下工程里比较稳的组合方式。这篇文章讲的是这套组合具体怎么搭、数据集怎么做、训练参数怎么给以及哪些位置会翻车。2. Unet 与 SAM 的组合方式提示框先验怎么进入分割网络2.1 SAM 的提示框为什么不能直接拼进 Unet很多初学同学从代码仓库里拿到一个SamPredictor就想当然地把predictor.predict(boxbox, multimask_outputFalse)的结果直接叠到 Unet 输出上做融合。这个做法在离线 demo 里能跑但放到训练里你会遇到两个问题。第一SAM 的 mask decoder 输出的是“物体掩码”不是“车道线掩码”。车道线是背景类的一部分SAM 的分割头很难稳定地把一条 3 像素宽的虚线完整认出来尤其是被车头遮挡或者处于阴影下的片段。第二SAM 的 prompt encoder 输出的是 token embedding它自己就是 transformer 语义空间把它直接 add 或 concat 到 Unet 卷积特征图上两者根本没有对齐反而会污染浅层的边缘特征。我在实践里用的是一个更直接的方案把提示框编码成一张空间上的高斯热图Gaussian Heatmap作为额外输入通道拼到 Unet 的输入端。这样做的好处是卷积网络天然吃空间结构不需要去对齐差异巨大的 token 空间。import torch import torch.nn.functional as F def box_to_gaussian_heatmap(boxes, size(512, 1024), sigma8.0): 把一组 [x1, y1, x2, y2] 的提示框转成高斯热图 boxes: (N, 4) 归一化到 0~1 的坐标 size: 与输入图像相同的高和宽 返回: (1, 1, H, W) 的 float32 张量多框时取最大值 H, W size N boxes.shape[0] heatmap torch.zeros((N, 1, H, W)) ys torch.linspace(0, 1, H).view(-1, 1).repeat(1, W) xs torch.linspace(0, 1, W).view(1, -1).repeat(H, 1) for i in range(N): x1, y1, x2, y2 boxes[i] cx (x1 x2) / 2 cy (y1 y2) / 2 # 框的宽高决定了高斯分布在不同方向上的延展 bw max((x2 - x1) * W, sigma * 2) bh max((y2 - y1) * H, sigma * 2) dist ((xs - cx) * W / bw) ** 2 ((ys - cy) * H / bh) ** 2 heatmap[i, 0] torch.exp(-dist).clamp(0, 1) return heatmap.max(dim0).values.unsqueeze(0)这段代码做的事情很直接把每个提示框的中心点当作高斯分布的均值框的宽高决定方差最后把所有框的热图取最大值合并成一张图。需要注意这里的坐标是归一化到 0~1 的如果从标注文件读到的是像素坐标先除以图像宽高再传进来。sigma参数控制的是热图衰减速率的基值通常取 6~12。车道线本身很细sigma 太大会导致热图糊成一片模型分不清哪条线是哪条太小又不足以覆盖 SAM 检测框的误差范围框稍微偏一点这个通道就等于没给信息。2.2 冻结 SAM 只做一个特征提取器既然不直接用 SAM 的分割结果那 SAM 模型本身还有没有参与的价值有而且是参与 Unet 的 decoder。我把 SAM 的图像编码器普通是 ViT-B 或 ViT-L拿过来在预训练状态下冻结用它提取原始图像的多尺度 token。这个特征里包含了极强的边缘先验而车道线又是一个极其依赖边缘和连续性的任务用它来补充 Unet 深层的上下文信息比单纯靠 Unet 自己卷出来的特征要稳。结构上的接法是这样import torch.nn as nn from unet_parts import DoubleConv, Down, Up from sam_adapter import SAMFeatureAdapter class LaneUnetSAM(nn.Module): def __init__(self, sam_image_encoder, in_channels4): super().__init__() # 输入 4 通道RGB heatmap self.inc DoubleConv(in_channels, 64) self.down1 Down(64, 128) self.down2 Down(128, 256) self.down3 Down(256, 512) self.down4 Down(512, 512) self.up1 Up(1024, 256) self.up2 Up(512, 128) self.up3 Up(256, 64) self.up4 Up(128, 64) # SAM 特征适配器把 ViT 的输出对齐到 Unet 特征图大小 self.sam_adapter SAMFeatureAdapter(sam_image_encoder) self.outc nn.Conv2d(64, 1, kernel_size1) def forward(self, x, sam_feats): x1 self.inc(x) x2 self.down1(x1) x3 self.down2(x2) x4 self.down3(x3) x5 self.down4(x4) # SAM 的高分辨率特征融合进最后一层上采样 sam_f self.sam_adapter(sam_feats, x5.shape[2:]) x self.up1(torch.cat([x5, sam_f], dim1), x4) x self.up2(x, x3) x self.up3(x, x2) x self.up4(x, x1) return self.outc(x)SAMFeatureAdapter的内部细节不复杂就是个1x1 卷积 双线性插值 残差的结构。SAM 图像编码器输出的 token 是 16x16 或者 64x64 的序列先 reshape 回二维特征图再用卷积把通道数降到和 Unet 当前层一致。冻结的操作在训练脚本里用requires_grad_批量处理。这里有个参数上的取舍值得展开说为什么不把 SAM 全部参数冻结完因为 ImageNet 和 SA-1B 预训练得到的特征偏向自然图像和车载摄像头视角下的路面有比较大的 domain gap。但解冻 ViT 全部参数会带来两个问题——显存暴涨、训练时间翻几倍。折中方案是冻结所有早层只解冻最后 2 个 transformer block 和 adapter 层在自建数据集上微调。2.3 提示框从哪来检测器和规则先验有人会问训练用提示框图推理时框又从哪里拿这里实际上分两套来源。第一套是离线标注时的真值框。做数据标注时用 LaneNet 这类模型先跑一遍把每条线的 bounding box 吐出来人工修正后存成本地 JSON 文件。在训练时这些框直接生成 heatmap模型在 supervised heatmap 下学习。第二套是部署推理时用轻量目标检测器比如 YOLOv8n先检测车道线区域输出若干个框再把框统一补丁成固定大小比如高度 64 像素、宽度 128 像素避免极细长的框导致热图被拉成一条没有宽度的线。需要注意一个常见的坑提示框不要直接套标注线上那样热图会变成一条极扁的高斯和图像本身的锐利边缘混在一起模型学到的是“提示框即边缘”的捷径。在线标注时把 bbox 的高度扩大 2~3 倍让高斯热图像一个淡淡的带状区域信息量反而更大。3. 自动驾驶道路线数据集准备TuSeries / BDD100K 转成语义掩码3.1 数据集选型对比标题里强调了“数据集”说明很多读者卡在第一步到底拿什么训练道路线分割。这里排除掉不太合适的选项按实际可操作性排个序数据集原始标注形式是否含线型类别适配难度推荐度TuSimplepolyline 坐标点序列只有车道线无类别低直接画线高适合起步BDD100Kpolyline含 road curb / lane有类别区分中需要区分线型高适合提升泛化KITTI Road只标注道路区域road area不是车道线高需二次生成低不推荐做 laneCityscapes有 lane 但标注稀疏部分序列有中一般KITTI 这个词在语义分割检索里热度很高但它的 Road 基准是“可行驶区域”不是车道线。如果你硬要拿 KITTI Road 做 lane 分割需要把标注边线 shrink 一个固定像素再人工核对工作量并不小。我一般把 KITTI 作为背景图像来源回填到训练集里做 domain augmentation而不是作为主训练集。3.2 TuSimple 标注怎么转成 maskTuSimple 的数据分布是一张 JPEG 配一个 JSON里面包含lanes字段每条线是一组点、h_samples固定的纵坐标列表、raw_file和lane_types。转换的核心工作是把点序列 rasterize 成像素掩码。import json import cv2 import numpy as np from pathlib import Path def tusimple_json_to_mask(json_path, img_size(720, 1280)): 把 TuSimple 单帧标注转成二值语义掩码 车道线画成 3 像素宽的白色折线 with open(json_path, r) as f: ann json.load(f) mask np.zeros((img_size[0], img_size[1]), dtypenp.uint8) h_samples ann[h_samples] lanes ann[lanes] for lane in lanes: # 过滤掉值为 -2 的无效点该 y 坐标上没有线 pts [] for y, x in zip(h_samples, lane): if x 0: pts.append((int(x), int(y))) if len(pts) 2: continue # 逐个线段连接避免整条 polyline 跨越大空隙 for i in range(len(pts) - 1): cv2.line(mask, pts[i], pts[i 1], 1, thickness3) return mask代码里有两个细节值得注意。一是-2这个无效值必须过滤TuSimple 的标注里有些纵坐标对应的车道线是缺失的。二是逐段cv2.line而不是直接cv2.polylines因为车道线在远处可能有断裂逐段画能保证每个连通段独立而且便于后续做连通域分析和数据增强。画线宽度取 3 像素是一个平衡点。太窄1 像素则下采样后直接消失太宽5 像素以上则模型学出的是“涂鸦条”而不是“线”推理时在远处容易产生粗尾。BDD100K 的格式类似只是它的属性里多了一个lane_style/lane_direction转 mask 时如果想做多类别分割给不同线型分配不同像素值即可。比如虚线是 1、实线是 2、双黄线是 3类别数量控制在 4~6 个以内避免尾部类别样本太少。3.3 语义分割数据集的增强要避开“随机擦除”通用语义分割里喜欢用的RandomErasing放在车道线任务上基本就是灾难。一张图里车道线的像素本来就少你再随机擦掉一块矩形区域等于是把一条线的中间段人工挖断。模型在训练时会学着把“擦除区域旁边的线”修复起来但这个行为在真实场景中没有任何对应物——真实车道线不会因为你的代码而凭空消失。我常用的增强策略是随机亮度对比度扰动、随机仿射变换旋转 5 度以内、水平剪切 0.05、色彩空间扰动HSV 通道随机偏移、随机加噪声模拟传感器噪声、重曝光模拟逆光。其中重曝光用cv2.addWeighted把线性图拉一个 gamma 曲线提升对阴影区域的鲁棒性。def lane_augment(image, mask): image cv2.cvtColor(image, cv2.COLOR_BGR2HSV).astype(np.float32) image[..., 0] np.clip(image[..., 0] np.random.randint(-5, 5), 0, 179) image[..., 1] np.clip(image[..., 1] * np.random.uniform(0.9, 1.1), 0, 255) image[..., 2] np.clip(image[..., 2] * np.random.uniform(0.8, 1.2), 0, 255) image cv2.cvtColor(image.astype(np.uint8), cv2.COLOR_HSV2BGR) # 仅对 mask 做相同的仿射变换保持空间对齐 angle np.random.uniform(-4, 4) M cv2.getRotationMatrix2D( (mask.shape[1] / 2, mask.shape[0] / 2), angle, 1.0 ) image cv2.warpAffine(image, M, (mask.shape[1], mask.shape[0])) mask cv2.warpAffine(mask, M, (mask.shape[1], mask.shape[0]), flagscv2.INTER_NEAREST) return image, mask仿射变换的插值方式对 mask 用INTER_NEAREST这一点不能妥协。如果用INTER_LINEAR线上会出现灰度过渡值BGR 的图片里看着不明显但一旦 mask 里有多类别1、2、3线性插值会造成语义混淆。4. 训练配置与源码结构冻结 SAM 训练 Unet 的关键参数4.1 训练脚本的最小闭环一段小的训练脚本方便看清楚数据流、模型流、梯度流的交接关系。这里省略了 dataloader 的样板代码只保留核心逻辑。import torch import torch.nn.functional as F from torch.cuda.amp import autocast, GradScaler def train_one_epoch(model, sam_encoder, dataloader, optimizer, scaler): model.train() sam_encoder.eval() # 冻结的 SAM 图像编码器 for imgs, heatmaps, masks in dataloader: imgs imgs.cuda() heatmaps heatmaps.cuda() masks masks.cuda().float() optimizer.zero_grad() with torch.no_grad(): # ViT 特征提取不需要梯度 sam_feats sam_encoder(imgs) # 组装带提示框通道的输入 x torch.cat([imgs, heatmaps], dim1) with autocast(): logits model(x, sam_feats) # 二分类线上/非线上 loss F.binary_cross_entropy_with_logits( logits, masks, pos_weighttorch.tensor([15.0]).cuda() ) # 辅助的边界损失让模型关注线的连续性 edge_mask (masks[:, :, 1:, :] ! masks[:, :, :-1, :]).float() loss 0.5 * F.l1_loss( torch.sigmoid(logits[:, :, 1:, :]), torch.sigmoid(logits[:, :, :-1, :]), reductionmean ) * edge_mask.mean() scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()这里pos_weight15是自行车道线任务里的一个牵制系数负样本远多于正样本让模型为正样本多出一点力。但 15 这个值不能拍脑袋决定应该先统计每张图车道线像素占比。如果你的数据集里车道线平均占 2%那pos_weight从 10 起步观察 loss 曲线后逐步调整。边界损失部分是我后来加的。代码里用相邻行做差分得到的是水平方向上的梯度图也就是“线边缘”的位置然后让 logits 在边缘处保持锐利。这个损失对虚线效果好——它不会像单纯 BCE 那样把断裂处的空隙也强行补齐而是保留真实的车道线间距。4.2 源码目录与权重加载顺序实际的源码组织按下面的方式展开每个文件只做一件事lane_seg/ ├── data/ │ ├── tusimple_dataset.py │ └── augment.py ├── models/ │ ├── unet_parts.py │ ├── sam_adapter.py │ └── lane_unet_sam.py ├── train.py ├── export_onnx.py └── configs/ └── lane_sam.yamlunet_parts.py就是标准的 Unet 模块DoubleConv、Down、Up 三层结构。sam_adapter.py负责把 SAM 的 ViT 输出重排成二维特征图并降通道。train.py读配置、构造模型、加载预训练、跑 epochs。权重加载的顺序有讲究。建议按这个顺序来先用公开的 SAM ViT-B 权重初始化sam_image_encoder不加载 mask decoder因为后文根本不用它。再用 ImageNet 预训练权重初始化 Unet 的 encoder 部分。最后把两个模块的权重分别 loadsam_encoder整个静置unet参与优化。# 伪代码只加载我们需要的部分 sam_ckpt torch.load(sam_vit_b_mask_encoder.pth, map_locationcpu) sam_state { k.replace(image_encoder., ): v for k, v in sam_ckpt.items() if k.startswith(image_encoder.) } sam_encoder.load_state_dict(sam_state, strictFalse) unet_ckpt torch.load(unet_backbone.pth, map_locationcpu) model.inc.load_state_dict(unet_ckpt[encoder.0], strictFalse)strictFalse在这里必须加因为 Unet 的加载可能是从不同主干里迁移的某些层名存在出入。但如果 loss 直接震荡不收敛首先排查的就是strictFalse是否有大量 key mismatch——如果 mismatch 的属性超过 80%说明权重加载实际上失败了模型相当于从头训练。4.3 训练参数的基准配置给一个基准配置适用于单张 24GB 显存卡输入分辨率 512x1024batch size 8。参数数值说明input_size512x1024保持车道线的细长结构不宜降到 256x512batch_size824GB 显存下的合理值梯度小的可以加到 16optimizerAdamWlr1e-4weight_decay1e-2backbone Unet lr1e-4主干保持稳定adapter lr3e-4新加的层可以更高schedulerCosineAnnealing50 epochsmin_lr5e-7pos_weight10~15视数据集中线比率而定SAM 冻结数全部冻结如果解冻最后 2 层lr5e-6EMA0.999测试时用 EMA 权重稳定性好很多输入分辨率 512x1024 是一个有争议的选择。用更大的 768x1536 会带来显存压力而且车道线的高频信息在下采样到 stride 8 之前都不会丢失512 宽在绝大多数车载摄像头画面里足够用。更关键的是图像 resize 时用INTER_AREA而不是INTER_LINEAR这样可以减少细线的模糊。4.4 训练权重的导出与复用训练完的权重文件不只是 torch 的 checkpoint还要考虑部署端的加载。我一般保存三份checkpoints/last.pt当前最后一步的完整状态含 optimizer 状态用来续训。checkpoints/best_iou.pt只存model.state_dict()去掉附带信息文件较小。checkpoints/best_ema_onnx导出 ONNX 格式输入是(1, 4, 512, 1024)的图像heatmap。ONNX 导出时需要特别注意一个坑SAM 的 ViT 是 transformer导出时如果把整个sam_encoder一起导出固化成全 tokens 流程推理时改动限制很大。实际的做法是仅导出 Unetadapter 那个分支输入张量从(1, 4, 512, 1024)开始SAM 特征在训练时已经被 adapter 吸收推理时这个模型并不需要单独再调 SAM。5. 验证与部署车道线分割不能只看 mIoU5.1 五项指标按重要性排序语义分割的标准评估指标是 mIoU但对车道线来说mIoU 有一个偏向它衡量的是像素重叠而车道线是极细结构只要预测线偏移 3~5 个像素mIoU 就会断崖式下跌反过来如果预测线有断裂但整体像素面积没变mIoU 可能反而没变化。所以我建议跟踪一组指标指标计算方式关注问题lane_mIoU只取线上类的 IoU直观中心线偏移预测线和真值线骨架之间的平均距离判断“偏了”连通性预测线中最大连通域的长度占比断裂检测F1T以 IoU0.5 为 TP 的逐线 F1实际可用性AP线上像素的 precision-recall 曲线虚警率其中“中心线偏移”在代码里不太好算简单做法是对预测 mask 做骨架提取skimage.morphology.skeletonize然后对每个骨架像素找到最近的真值骨架像素求距离平均。这个值如果大于 10 像素就说明热图通道里的提示框当前没有起到修正作用可能是框的位置本身漂移太大。5.2 每 N 步可视化一次肉眼是最终标准训练过程里我习惯每 500 步输出一组对比图原图 SAM 提示框、真值 mask、预测 mask、预测和真值 overlay。看这四个输出能快速定位三种典型故障远处车道线全是碎点大概率是下采样时线太细。把输入分辨率调高或者在 loss 里对远端行区域加一个 2x 权重。车道线被“补全”到不该有线的位置典型是虚线部分被填实边界损失的权重不够。线整体偏移但不碎提示框和模型输出不一致检查 SAM 的框生成器和 Unet 的输入是否对齐宽高比可能出问题。最后说一个非常有用的技巧在验证集里固定挑选 20 张有代表性的图雨天、夜间、逆光、匝道每轮 epoch 跑完后强制输出这 20 张的结果到一张大图上。这样你不需要翻 tensorboard 就能一眼看出模型的退化方向尤其是在数据增强过强导致模型变得畏首畏尾的时候这个列表比 mIoU 曲线可靠得多。本文还有配套的精品资源点击获取