简介图像分割是计算机视觉的基石任务之一传统方法依赖针对特定场景精心设计的网络结构泛化能力有限。SAMSegment Anything模型的出现以提示学习机制将分割转化为通用能力极大降低了领域适配成本。然而在光照不足、遮挡严重或纹理模糊等复杂环境下单一可见光模态的信息瓶颈会导致分割精度显著下降。双模态分割通过融合深度图、红外图等互补信息为模型提供额外的几何或热特征从而有效提升复杂场景下的鲁棒性。本文从深度学习与SAM基础原理出发剖析了在SAM架构上引入双模态输入的技术路径并给出了从数据对齐、模型加载到微调训练、显存优化的全流程工程实践涵盖自动驾驶、工业质检、遥感分析等典型应用场景最终自然落地于一个可运行的双模态SAM分割项目帮助开发者快速复用与二次开发。 开头直接进入主题不客套。SAMSegment Anything这个模型在图像分割领域确实掀起了一波热潮Meta开源之后整个行业的玩法都变了。以前我们做分割要么是UNet换各种backbone要么是YOLO加分割头训练一套数据就要重新调一套参数。SAM不一样它把“万物可分割”变成了通用能力你再也不用为一个特定类别去专门设计网络结构了。而这个项目——基于SAM架构的双模态图像分割拿到手的第一感觉就是它在SAM这条通用分割路径上把输入从单一的RGB图扩展到了两种模态比如可见光加深度图、可见光加红外图这样。这种设计直接解决了一个很实际的问题纯视觉信息在光照差、遮挡多、纹理不清晰的时候分割效果会断崖式下跌而多一路模态数据就能把精度拉回来不少。我花了两天时间把整套代码、数据和已训练好的权重完整跑通了包括训练、验证、推理、自定义数据替换。今天这篇就直接说清楚它的原理、如何跑起来、会遇到哪些坑以及怎么改成你自己的双模态数据。不管你是做自动驾驶、工业质检、遥感分析还是医学影像只要想用SAM做双模态分割这套东西都能给你省下至少两周时间。1. 项目核心思路与架构解析1.1 为什么要在SAM基础上做双模态分割先回答一个最基础的问题我们已经有了SAM这么强的分割模型为什么还要做双模态不是说SAM不够好而是实际场景里单靠RGB图像本身有天然的短板。举个例子自动驾驶里面的夜间场景一张普通摄像头拍出来的图行人和路边的围栏颜色很接近亮度也低你让SAM自动分割它可能把人都和地面背景混在一起。但如果这时额外输入一路深度图Depth行人的轮廓和路面在距离上有明显断裂模型就很容易把目标区分出来。再比如工厂质检里检测产品表面的缺陷有些划痕在可见光下几乎看不出来但在红外光下特征非常明显。双模态分割的价值就是提供了“互补信息”一个模态不够就再加一个模态。这个项目把双模态的数据输入到SAM架构里相当于在保持SAM现有分割能力的基础上让模型多了一条“感知通道”。实测下来在弱光、遮挡、目标边界模糊等场景下双模态对比单模态的mIoU能稳定提升5到10个百分点这个提升幅度在工程上是很可观的。1.2 SAM架构的几个关键组成要理解这个项目的代码结构得先把SAM本身的架构拆开看。SAM由三个核心部分组成。图像编码器Image Encoder是ViT变体用于把输入图像转换成高维特征向量。它输出的特征图保留了空间信息后续的分割掩码就是基于这层特征解码出来的。提示编码器Prompt Encoder负责处理用户提示包括点、框、掩码这些。提示的作用是告诉模型“你要分割的是哪个目标”这也是SAM能实现zero-shot分割的关键。掩码解码器Mask Decoder把图像特征和提示特征融合最终输出分割掩码和对应的置信度。这个项目的双模态扩展方式并不复杂核心是改输入融合层。它保留了SAM的图像编码器和掩码解码器的预训练权重只是在图像编码器前面增加了一个模态对齐层把第二路模态数据比如深度图编码成和RGB特征兼容的向量再把两路特征在通道维度上拼接或加权融合最后输入到SAM的编码器里。这种做法有一个很大的工程优势不必从零开始训练一个巨大的ViT模型而是可以复用官方在大量数据上预训练好的SAM权重只额外训练新增的模态融合层。整个项目的训练成本和控制难度都大幅下降。1.3 项目代码结构与关键文件一览我拿到项目后的第一件事是把整个目录过一遍。它没有把所有东西都塞在一个文件里而是分了几个模块组织方式比较清爽。checkpoints/ 预训练权重文件夹 sam_b.pt 官方SAM ViT-B的权重用于加载初始参数 multi_modal_sam.pt 双模态模型训练好的完整权重 data/ train/images/ 训练集RGB图 train/depths/ 训练集深度图 train/masks/ 训练集标签掩码 val/images/ 验证集RGB图 val/depths/ 验证集深度图 val/masks/ 验证集标签掩码 src/ dataset.py 数据加载与双模态预处理 model.py SAM双模态模型定义 train.py 训练入口 infer.py 推理入口 utils.py 评估指标与可视化辅助函数 configs/ config.yaml 模型与训练参数配置这种结构一看就是为实际项目设计的不是竞赛玩具。我建议你也养成这样的习惯数据和代码分离、配置和逻辑分离。后面换数据集或者换参数直接改config就行不用动代码。2. 数据准备与已训练模型的正确打开方式2.1 双模态数据集组织与对齐规则拿到项目后很多人第一反应是直接跑然后报错“shape不匹配”或者训练半天loss不降。这大概率是数据没对齐。双模态分割里RGB图、第二模态图、标签掩码三者必须确保像素级对齐。什么是像素级对齐就是同一张图RGB图像坐标为(x, y)的那个像素和深度图坐标(x, y)的像素以及掩码标注中(x, y)的像素必须对应现实世界中的同一点。任何一方的分辨率不同、任何一方的拍摄视场角不同都会导致训练时模型学到错误特征。这个项目默认使用的数据已经完成了对齐。RGB图是512x512、第二模态图是512x512、掩码是512x512。如果你的数据集多路图像分辨率不一致需要先做resize或crop。这里我提醒一句resize时RGB、深度图、掩码三者必须使用完全相同的插值方式和目标尺寸否则会额外引入像素偏移。特别是掩码建议用PIL的NEAREST插值不要用双线性插值否则会引入标注外的中间像素值。2.2 预训练权重加载与文件校验项目里有两个权重文件让我一开始有点困惑sam_b.pt和multi_modal_sam.pt。后来仔细看代码才明白这是两阶段策略。sam_b.pt是官方SAM的ViT-B权重它的作用是初始化图像编码器和掩码解码器的参数让双模态模型在训练初期就具备基本的分割能力。multi_modal_sam.pt是作者训练完成后的完整双模态模型权重包含所有模块参数。直接推理时只需要加载multi_modal_sam.pt。实际运行时需要注意权重文件路径和大小。我拿到文件后先做了个简单检查SAM官方ViT-B权重大概375MBmulti_modal_sam.pt从大小上就能看出是否包含了融合层的额外参数。如果加载时提示key不匹配比如“missing keys”或者“unexpected keys”先检查是自己加载了哪个权重文件再看state_dict的前缀模块名是否和model.py定义一致。这个问题在第4章会再展开讲。提示不要用GB级别的ViT-L或ViT-H权重直接套这个项目的模型除非你显存充足且愿意重新适配输入尺寸。项目的预训练模型是基于ViT-B训练的换更大的backbone需要重新验证输入特征对齐方式。2.3 用预训练权重直接跑验证集跑通流程永远是第一步。训练前先验证模型能出结果再谈优化。这个项目我用以下命令在验证集上做了一次快速推理整个过程大概一分钟。python src/infer.py --config configs/config.yaml --checkpoint checkpoints/multi_modal_sam.pt --split val如果一切正常终端会打印出每张验证图的mIoU、Pixel Accuracy等指标同时会在输出目录生成三张图原图、掩码叠加图、预测掩码。我跑完后验证集mIoU在0.83左右和作者README里报告的数值基本一致。这说明模型推理链路是通的权重没有损坏。这里也给第一次接触这个项目的人一个判断标准你的输出目录里如果没有生成可视化图只有指标数字说明代码走到了评估函数但可视化函数在静默报错。检查一下输出目录权限和opencv的写入格式即可。3. 完整代码实现与关键模块拆解3.1 模型定义如何把两路模态喂给SAM核心代码在model.py里这部分是整篇项目的灵魂。作者没有直接fork官方的SAM代码而是在其之上加了一个双模态适配层。我分几个部分说明。第一段是构建基础图像编码器。它加载原始SAM的ViT模型但把最最开始输入层的卷积改成接受两路输入。用如下代码段展示import torch import torch.nn as nn from segment_anything import sam_model_registry from segment_anything.modeling.image_encoder import ImageEncoderViT class MultiModalSAM(nn.Module): def __init__(self, sam_checkpoint, fusion_modeconcat, depth_input_channels1): super().__init__() # 加载官方SAM ViT-B sam sam_model_registry[vit_b](checkpointsam_checkpoint) self.image_encoder sam.image_encoder self.prompt_encoder sam.prompt_encoder self.mask_decoder sam.mask_decoder # 双模态融合层把深度图1通道编码到和RGB特征一致的维度 self.depth_proj nn.Sequential( nn.Conv2d(depth_input_channels, 3, kernel_size3, padding1), nn.GELU(), nn.Conv2d(3, 3, kernel_size3, padding1) )这里的关键是depth_proj这个模块。它的作用不是把深度图变成RGB而是把单通道的深度值映射到三通道特征空间好让视觉编码器能处理。训练时这个模块的参数会被更新去学习“如何把深度信息转换到视觉特征域”。第二段是前向传播逻辑。def forward(self, rgb, depth, boxesNone, pointsNone): depth_feat self.depth_proj(depth) fused_input rgb depth_feat image_embedding self.image_encoder(fused_input) # 提示编码与掩码解码沿用官方SAM逻辑 ...这里用的是加法融合比较轻量。另一种常见的做法是通道拼接后接一个1x1卷积降维我用下面这个对比表格说明差异。融合方式实现成本参数增量适用场景加法融合极低一个proj层两路模态特征图尺寸相近且互补性强通道拼接中等1x1卷积两路模态维度差异大需要充分交互注意力融合较高加注意力模块两路模态包含大量冗余与噪声需要自适应降噪我在实际测试中这个项目选择的加法融合在深度图上效果已经不错。如果换到低质量的第二模态数据比如带噪声的红外建议换成通道拼接再加几层卷积鲁棒性会更强。3.2 数据加载器里的双模态预处理细节光有模型还不够数据加载器必须保证每一批数据都成对出现。dataset.py里有一个Dataset类我看了下实现逻辑并不复杂但有几个细节做得比较到位。它读RGB图用PIL读标存用numpy把深度图当作普通灰度图加载。所有图像都通过torchvision的Resize统一到固定尺寸归一化的均值方差也用了ImageNet的标准参数。以下是关键代码段class DualModalDataset(torch.utils.data.Dataset): def __init__(self, image_dir, depth_dir, mask_dir, image_size512, modetrain): self.image_paths sorted(glob.glob(os.path.join(image_dir, *.png))) self.depth_paths sorted(glob.glob(os.path.join(depth_dir, *.png))) self.mask_paths sorted(glob.glob(os.path.join(mask_dir, *.png))) self.image_size image_size self.mode mode def __len__(self): return len(self.image_paths) def __getitem__(self, idx): image Image.open(self.image_paths[idx]).convert(RGB) depth Image.open(self.depth_paths[idx]).convert(L) mask Image.open(self.mask_paths[idx]).convert(L) # 关键三路数据必须用完全相同的随机裁剪/翻转操作 if self.mode train: image, depth, mask self._train_transform(image, depth, mask) else: image, depth, mask self._val_transform(image, depth, mask) return torch.tensor(image), torch.tensor(depth), torch.tensor(mask)特别注意_train_transform里对三路图同时做的随机翻转和裁剪这一点非常重要。如果你像处理普通单模态分类图那样用独立的随机增强库会导致深度图和RGB图空间位置错位模型训练出来精度一定高不了。我当时踩过这个坑第一次自己写数据集类时分别对RGB和depth各自随机翻转结果训练了20轮mIoU只有0.3。后来看到官方给的代码里用了一个统一的随机状态问题才解决。3.3 推理脚本参数调整与输出可视化infer.py这个脚本比较短主要做了四件事加载配置、加载模型权重、遍历数据集推理、计算指标并输出可视化结果。其中有两个参数值得你根据实际情况调整。第一个是batch_size默认配置里是4。如果你显存只有8G建议调到2或1否则很快会OOM。第二个是--output_dir默认在outputs目录这里最好每次实验换成新目录避免被历史结果覆盖。我自己习惯命名成outputs_val_0325或outputs_test_run1方便对比不同实验。推理脚本里还做了一个值得学习的细节它对预测掩码做了argmax归一化把类别索引映射到0到255的灰度值然后用colormap上色叠加到原图上。这样输出结果直接就能看清哪些区域被分割成了哪一类。如果你的分割类别数比较多记得改colormap的映射大小否则不同类别显示的颜色会撞。3.4 基于新数据的微调训练如果你不想直接使用预训练权重而是想在自己的双模态数据上微调需要执行训练命令。python src/train.py --config configs/config.yaml默认的config里设置了50个epoch、batch_size为4、学习率1e-4、优化器AdamW。这个配置在项目自带的数据上跑了大约40分钟一张RTX 3090。我自己的实验里用更小的数据集微调时把epoch降到了20同时把学习率调到5e-5避免在少量数据上过拟合。训练过程中日志会打印每个epoch的loss和验证集mIoU。我特别关注一个信号如果训练loss开始下降但验证集mIoU停滞不前说明模型容量对于数据量来说太大了需要加正则化或调低学习率。如果训练和验证mIoU都在涨恭喜你说明模态融合层确实学到了有用的跨模态特征。4. 常见问题与排查技巧实录4.1 环境与依赖问题版本不匹配是最容易翻车的地方这个项目使用的是segment-anything官方库同时依赖PyTorch、OpenCV和PyYAML。我第一次运行时直接在自己的PyTorch 2.0环境里装了最新的segment-anything结果报错说运行时版本不兼容后来发现是segment-anything的版本和PyTorch版本存在耦合。建议用一个干净的环境来跑这个项目具体命令如下conda create -n sam_dual python3.9 conda activate sam_dual pip install torch1.13.1 torchvision0.14.1 pip install githttps://github.com/facebookresearch/segment-anything.git pip install opencv-python pyyaml版本选1.13.1的原因是这个项目训练时用的就是这个环境测试起来最稳妥。用更新的PyTorch也能跑但可能遇到个别API变化。4.2 加载预训练权重时报错key不匹配项目里最常被问到的报错之一是load_state_dict时提示缺失键或不认识的键。这个问题几乎都是因为加载权重文件选错了。如果你用的是完整的multi_modal_sam.pt必须在完整模型对象上调用load_state_dict也就是先创建MultiModalSAM实例再加载。如果你先加载了官方的sam_b.pt来初始化那么只会给图像编码器和掩码解码器赋初值depth_proj等新模块还是随机值。以下是正确写法model MultiModalSAM(sam_checkpointcheckpoints/sam_b.pt) state torch.load(checkpoints/multi_modal_sam.pt, map_locationcpu) model.load_state_dict(state[model_state_dict])另一个可能的问题是模型定义修改后之前保存的key少了一些模块比如你把depth_proj的名字改成了depth_encoder加载旧权重时就会missing key。解决办法是把模型定义恢复原样或者只加载已有参数并忽略缺失项。4.3 推理结果产生大量空洞或毛刺后处理优化建议如果你用预训练权重在自己数据上做推理发现掩码轮廓有大量空洞、毛刺边缘不平滑不要急着怀疑模型没训练好。先检查输入数据的预处理是否和训练时一致尤其是归一化和resize尺寸。如果数据本身没问题那可以增加一步后处理来改善视觉质量。比较常用的是条件随机场CRF或者简单的形态学闭运算。我在项目中加了一段轻量级的后处理用open-cv的morphologyEx去除小噪点效果立竿见影import cv2 import numpy as np def postprocess_mask(mask): kernel np.ones((3, 3), np.uint8) mask cv2.morphologyEx(mask, cv2.MORPH_CLOSE, kernel, iterations2) mask cv2.medianBlur(mask, 5) return mask注意后处理只对视觉结果有帮助不会提升mIoU指标。如果目标是刷指标最好还是从训练数据增强或模型结构上下功夫。4.4 显存不足与推理速度优化的实操方案由于SAM本身就基于ViT显存消耗不小再加上双模态输入显存压力更大。我在8G显存的笔记本上实测batch_size设为1可以完成推理但训练基本不现实。如果你显存紧张可以用以下几个招数。第一是降低输入分辨率。模型默认输入512x512如果你改成384x384显存占用能降将近一半同时显存mIoU损失大约1到2个点。这个改动需要在config.yaml里修改image_size且推理和训练必须保持一致。第二是开启混合精度。在infer.py或train.py里加入torch.cuda.amp.autocast()以后推理显存占用会降低约30%速度提升10%到20%。我实测下来这个项目适合开fp16因为模型已经训练得较充分精度损失可以忽略。第三是切片推理。如果你的输入图片非常大比如遥感图像是几千乘几千的尺寸直接resize到512会丢失太多细节。这时候可以把图像切成多个512x512的小图分别推理后再拼接在一起。注意切分时要有重叠区域拼接时对重叠部分取平均值否则接缝处会出现清晰的分割断层线。4.5 换到自己数据上从组织格式到label命名如果你想用这个框架跑自己的双模态数据不要一上来就改代码先把数据整理成项目默认的目录格式。这个坑我替你们踩过目录不对会让各种glob匹配返回空列表训练时会报“找不到任何图像”。标准的数据组织方式是data/ train/ images/ # rgb_0001.png, rgb_0002.png ... depths/ # depth_0001.png, depth_0002.png ... masks/ # mask_0001.png, mask_0002.png ... val/ images/ depths/ masks/文件命名需要保持统一且一一对应比如rgb_0001.png与depth_0001.png和mask_0001.png是一组。在dataset.py里三个目录的排序用了sorted()所以只要你能保证三组文件的前缀命名一致对应的顺序就不会错。我自己的经验是图片文件格式建议统一用png因为它是无损格式深度图保存为jpg会压缩掉很多深度细节对分割精度影响很大。如果原始数据是jpg也建议先批量转换成png。注意mask里如果包含多类别标签每个像素值应该是0,1,2,...这样的类别索引并且背景是0。如果你的mask是RGB彩色标注需要先转成灰度索引图否则模型会把它当多通道输入处理直接报错或训练出无效结果。5. 为什么这个项目的技术路线值得复用做图像分割这些年我见过太多从零训练的网络也见过太多调不动效果的单模态方案。SAM架构的出现确实改变了工作方式。它告诉我们一个道理通用分割先验能力非常重要而在此基础上引入额外模态是一条非常实用的演进路线。这个项目的价值不在于它的每行代码多么炫技而在于它把一个“可行方案”完整落地了数据组织、权重转换、融合层设计、训练与推理流程全链路清晰。这比网上那种只有一个高精度报告的“标题党项目”实在太多。从工程角度看它的扩展性也很好。你不需要局限于RGB加深度图只要把第二路模态替换成红外图、热成像图、甚至医学影像里的CT断层图再把depth_proj的输入通道数改一下就能适配很多领域。这也是我在文章开头说它能节省两周时间的原因。我在实测过程中把第二路模态从深度图换成近红外图只改了输入通道和部分预处理参数流程没有任何阻塞验证集精度和原模型相当。这说明双模态适配层的设计是通用的本质是学习不同信息源的投影方式而不是绑定某一种特定模态。6. 最后再分享一个实操小技巧训练完成后我建议你额外跑一次全验证集的PR曲线和混淆矩阵。项目自带的评估函数只输出mIoU但实际部署时你可能更关心哪些类别容易被混淆。我在做广告牌背景分割时就发现广告牌边框和天空背景的混淆比例很高后来通过给mask加了一层边界监督把边界像素的损失权重调大混淆情况明显缓解。这个技巧虽然老套但在双模态场景里更有效因为两路模态在边界处的特征差异通常比内部更明显。深度图上目标的边缘一般有深度值突变RGB上也有颜色突变两边同时给监督信号模型更容易学到稳定的边界特征。你在用这个项目时如果遇到类别边界黏连问题可以重点关注这一步。本文还有配套的精品资源点击获取