ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

SAM双模态图像分割系统:完整代码、数据集与训练权重实践指南

SAM双模态图像分割系统:完整代码、数据集与训练权重实践指南 简介图像分割是计算机视觉中的核心任务传统语义分割模型往往受限于固定类别换一个场景就需要重新标注和训练。以SAM为代表的提示分割模型通过点、框或文本等条件提示生成目标掩膜将分割任务从封闭分类转变为开放世界的条件生成。双模态融合进一步将文本语义与图像特征对齐使同一套模型能够灵活应对广告牌检测、工业质检、遥感目标分割等多种场景。本文围绕一个基于SAM架构的双模态图像分割工程解析其架构设计、文本与点框提示的融合原理、数据构造方式、训练与推理的完整实现并分享实测中的调参与排错经验。工程已提供完整代码、标注数据集和预训练权重可直接复现或二次开发适合需要快速落地通用分割能力的开发者参考。 做图像分割的同行应该没少被问过这样一个问题能不能给我一套开箱即用的分割模型代码要全、数据要齐、权重还得是训练好的说实话网上关于SAMSegment Anything Model的demo和论文解读已经很多了但真到自己要训练、要改双模态输入、要拿现成权重跑推理的时候大多数人还是卡在跑不通缺数据没权重这三座大山上。这个项目正好把这几个痛点一次性解决了基于SAM架构做了一套双模态图像分割系统带完整代码、带数据集、带训练好的模型权重我按文档实际跑过一遍训练和推理都能正常走通直接拿来复用或者在此基础上做二次开发都没有问题。这篇文章我会把整个工程的来龙去脉讲透包括为什么选SAM架构、双模态到底怎么融合、数据是怎么构造的、训练和推理的完整逻辑是什么以及我在实测过程中踩过的坑和排查方法。无论你是做学术研究需要对比实验还是在工业场景里落地广告牌检测、工业质检、遥感目标分割这类需求这篇内容应该都能帮你省下不少时间。1. 项目价值与设计思路1.1 双模态分割到底解决什么问题传统语义分割模型比如U-Net、DeepLabV3、PSPNet这些本质上做的是一个封闭集合的分类任务。模型训练的时候定了哪几个类别推理的时候就只能分哪几个类别。今天项目里要检测广告牌明天换一个场景要检测消防栓那就得重新标注数据、重新训练模型整个流程非常重。而SAM这类提示分割模型完全换了一个思路输入一张图再给一个提示prompt模型根据提示来生成对应目标的掩膜。提示可以是点、可以是框也可以是文本描述。这就把分割从固定类别分类变成了条件生成——模型本身不关心你要分什么它只负责理解提示并找到图里对应区域。这个项目里说的双模态指的就是同时输入图像和另一种模态的提示信息。最常见也最实用的组合是这两种图像 文本描述比如输入一张街道图文本写红色消防栓模型输出消防栓的像素级掩膜。图像 点/框坐标点落在目标内部模型分割出整个目标框框住目标模型细化出精确边界。这两种输入模式组合在一起覆盖面就广了。广告牌检测、遥感建筑提取、医学影像器官分割、工业缺陷定位甚至电商场景里的商品抠图都可以用同一套模型来支撑。这也是这个项目最有价值的地方不用为每个任务重新训练一个完整的分割模型只需要在推理时换一下提示模型就能适配不同目标。1.2 为什么选SAM架构而不是传统分割模型可能有人会问双模态分割我用CLIPSeg、用Grounding DINO加SAM的pipeline或者干脆用U-Net加一个文本编码器不行吗都可以但这里有一个工程选型问题。从泛化能力来看SAM在SA-1B这个超大规模数据集上预训练过图像编码器提取到的特征是通用的对不同域的数据都有很强的适应能力。如果从零训练一个U-Net在数据量有限的情况下泛化效果很难比得上微调SAM。从提示设计来看SAM原生支持点、框、mask这三种稀疏/稠密提示输入mask decoder本身就是在这些提示条件下做分割推理的。我们要加文本模态只需要把文本embedding对齐到提示embedding的空间里就能复用整套SAM的解码机制改动量非常小。从成本来看SAM的image encoder虽然参数量大但训练时完全可以冻结只需要微调mask decoder和prompt encoder。这比从头训练一个分割网络要便宜得多。我用下面这个表格把几个常见分割方案放在一起做对比方便你根据场景选型方案提示支持泛化能力训练成本适合场景U-Net系列无固定类别弱换场景需重训低单任务、数据量充足的场景DeepLabV3无固定类别中等中等通用语义分割Mask R-CNN检测框实例中等中等实例分割依赖检测框质量SAM点/框点、框、mask强低微调decoder交互式分割、开放世界分割SAM文本分支本项目文本、点、框强低双模态提示分割、跨模态检索分割结论很直接如果任务固定、类别固定、数据量也很充足传统分割模型依然是性价比之选。但如果你要做的是一个能适应多种目标、支持多种交互方式的通用分割系统SAM架构就是当前最合理的底座。2. SAM架构核心原理拆解2.1 三大核心组件图像编码器、提示编码器、掩码解码器要把SAM用好不能光停留在调用API的层面至少得把它的三个核心组件搞清楚。第一个组件是图像编码器Image Encoder。它本质上是一个用MAE方式预训练过的ViT模型。输入是一张RGB图像经过ViT的patch embedding和一系列transformer层之后输出一个分辨率降低但通道数很高的特征图。SAM官方提供了vit_b、vit_l、vit_h三个规格参数量分别是91M、308M、637M左右。这个特征图包含了图像的全局语义和局部细节信息后续的分割掩膜就是从这套特征里解码出来的。第二个组件是提示编码器Prompt Encoder。提示分两种稀疏提示和稠密提示。点和框属于稀疏提示每个点会被编码成一个embedding向量框则用左上角和右下角两个点的embedding表示如果某个点在目标外部还需要用一个负点标记来告诉模型这里不是目标。mask这类稠密提示会经过卷积层降采样和图像特征对齐。这里的关键在于提示编码器把不同模态的提示统一成了同一种表示空间这让后面接入文本模态变得非常自然。第三个组件是掩码解码器Mask Decoder。它接收图像特征和提示embedding通过两层的双向Transformer进行跨模态交互再经过动态mask预测和上采样最终输出一个或多个候选掩膜。SAM默认输出三个掩膜对应整体、子部分和超集三种粒度的分割结果推理时可以根据置信度或交互反馈选择最合适的一个。2.2 双模态信息到底怎么融合进SAM项目里最核心的技术细节就是文本模态怎么和SAM融合。这里我分享一下我的实现思路这也是目前社区里比较主流的做法。第一步用CLIP的文本编码器把输入的文本描述编码成一个向量。CLIP本身是图文对齐模型文本编码器输出的特征天然带有语义信息比随机初始化一个文本网络要靠谱得多。第二步做一次维度对齐。CLIP文本编码器输出的维度一般是512或者768而SAM的prompt encoder输出的稀疏提示embedding维度是256。所以我加了一个投影层把CLIP特征从512维映射到256维让两种提示embedding能够直接拼接。这里我贴一下核心实现代码方便你理解这个融合层是怎么写的import torch import torch.nn as nn class TextPromptFusion(nn.Module): 把CLIP文本特征投影到SAM prompt embedding空间 clip_dim: CLIP文本特征维度通常为512 sam_dim: SAM稀疏提示embedding维度固定为256 def __init__(self, clip_dim512, sam_dim256): super().__init__() self.proj nn.Sequential( nn.Linear(clip_dim, sam_dim), nn.LayerNorm(sam_dim), nn.ReLU(inplaceTrue) ) def forward(self, text_feat): # text_feat: [B, clip_dim] 已做过全局池化的文本向量 return self.proj(text_feat) # [B, sam_dim]第三步在训练和推理时把这256维的文本embedding和SAM原始的稀疏提示embedding在token维度上拼接起来。SAM的mask decoder本身就是在图像token和提示token之间做cross-attention多一个文本token它照样能处理这样就不需要改动SAM内部的结构。点框提示的处理方式相对直接但有两个细节容易被忽略。一是坐标一定要做归一化把像素坐标除以图像宽高映射到[0, 1]区间否则训练和推理时图像尺寸一旦变化模型就失效了。二是正负点的语义很重要点在目标内部记为1在目标外部记为0这个标记直接影响模型对关注什么的理解。2.3 从原理到工程显存和速度的权衡SAM虽然效果好但不是没有代价的。vit_h规格的图像编码器单张图推理一次大概要占10GB以上的显存训练时用自动混合精度也很容易OOM。所以这个项目我在工程上做了一些取舍。模型规格上推荐优先使用vit_b。它在精度和显存消耗之间最平衡双模态微调场景下mask decoder和prompt encoder的参数量只占整体很小一部分真正吃显存的是image encoder。如果显存不够就把image encoder冻结住只优化prompt encoder、文本投影层和mask decoder这样显存占用能降下来一大截。另外文本编码器只在数据加载阶段离线提取一次特征不参与端到端的训练这样能省掉backward的计算图和显存占用。文本embedding提前存成npy文件训练时直接加载比每次forward都过一遍CLIP要快很多。3. 环境准备与数据构造3.1 运行环境与依赖版本这个项目对硬件的要求是建议有10GB以上显存的GPU最好是RTX 3090、4080或者A5000及以上。CPU推理不是不行但速度会慢到让人怀疑人生不推荐。软件环境方面我实测可用的版本组合是这样操作系统Ubuntu 20.04 / 22.04Windows 10/11也能跑但数据预处理脚本建议在Linux环境执行Python版本3.9或3.10PyTorch2.0及以上CUDA 11.8或12.1segment-anything库最新的0.1.0版本即可安装命令我放在下面可以直接执行pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install segment-anything opencv-python pycocotools matplotlib numpy3.2 数据集格式与目录结构很多开源项目跑不通问题往往出在数据格式不统一上。这个项目的数据目录结构很简单三个部分data/ ├── images/ # 原始图像 ├── masks/ # 对应的掩膜单通道PNG0为背景255为目标 └── annotations.json # 标注信息索引annotations.json里面每个样本包含以下字段我截一个示例{ image: images/0001.jpg, mask: masks/0001.png, text: [广告牌, billboard], box: [120, 80, 340, 260], point: [200, 150], point_label: 1 }text可以同时给中文和英文box是目标外接框的左上右下坐标point是目标内部的一个点坐标point_label为1表示正点。实际训练时如果某条样本没有box可以从mask的外接矩形生成没有point可以从mask内部随机采样。如果你有自己的数据想转成这个格式有两个办法。一个是用labelme这类标注工具标好后写个脚本转成mask和json另一个是如果已有COCO格式的标注可以用pycocotools把annotation解析出来生成相同结构。3.3 数据增强与样本处理细节数据增强直接抄了SAM训练时的常规做法但有几个细节需要特别留意。图像和mask同步缩放随机缩放到640到1024像素之间的尺寸然后长边裁剪到1024。缩放时mask必须用最近邻插值用双线性会把边缘搞糊影响分割精度。提示点采集策略正点从mask内部随机采样负点从mask外部但距离目标不远的区域采样。正负点的比例一般控制在11左右。如果负点离目标太远模型会很快学会离远了就是背景这个特征学起来没有意义反而会让模型在目标边界附近表现变差。文本增强同一样本可能会用多个同义描述比如广告牌和billboard。训练时随机选一个描述输入CLIP提取特征这个操作很简单但对泛化能力提升非常明显。4. 完整代码实现与核心逻辑4.1 项目目录结构与模块划分一个工程能不能被别人顺利复现目录结构往往是第一道门槛。这个项目我按下面这种方式组织每个模块的职责很清晰sam-dual-modal/ ├── checkpoints/ # 保存训练好的模型权重 ├── data/ # 数据集目录 ├── src/ │ ├── dataset.py # 数据加载与提示生成 │ ├── model.py # 双模态SAM模型封装 │ ├── train.py # 训练脚本 │ ├── infer.py # 推理脚本 │ └── utils.py # 损失函数、评估指标等工具 ├── requirements.txt └── README.md4.2 数据加载器实现要点dataset.py是整个项目的基础它的职责是把raw data转成训练需要的张量。核心逻辑分三步读图、读mask、构造提示。读图用opencv注意要把BGR转成RGB因为SAM预训练用的是RGB顺序。读mask用PIL或者cv2的IMREAD_GRAYSCALE归一化到[0, 1]。提示的构造这里要重点讲一下尤其是点的采样策略。def sample_points(mask, num_positive1, num_negative1): 从mask内部和外部采样提示点 mask: [H, W], 0/1二值掩膜 返回point_coords和point_labels import numpy as np positive_idxs np.argwhere(mask 0) negative_idxs np.argwhere(mask 0) # 正点从目标内部随机选 pos_choice positive_idxs[np.random.choice(len(positive_idxs), num_positive)] # 负点从背景随机选 neg_choice negative_idxs[np.random.choice(len(negative_idxs), num_negative)] coords np.concatenate([pos_choice, neg_choice], axis0)[:, ::-1] # 转成(x, y) labels np.array([1] * num_positive [0] * num_negative) return coords.astype(float), labels.astype(float)这里代码里的[::-1]是把行列坐标转成x,y坐标这个细节如果不注意后面训练出来的模型会很奇怪。用[x, y]还是[row, col]本身没有对错但必须和输入SAM的坐标规范保持一致。4.3 训练脚本的核心逻辑与损失函数这个项目的训练策略是冻结image encoder只微调mask decoder、prompt encoder和文本投影层。这样做有两个直接好处一是显存占用低二是训练速度快因为ViT那部分反向传播的计算量被省掉了。损失函数用的是focal loss加dice loss的组合。focal loss解决的是前景背景像素数量极度不均衡的问题广告牌、小目标这类场景下目标像素往往只占整张图的很小比例focal loss的调制因子能让模型把注意力集中在难分样本上。dice loss则直接优化分割结果和真值的重叠程度对边缘质量有正面作用。训练循环的核心代码大概是这样的我按可读性做了简化import torch import torch.nn.functional as F def train_one_epoch(model, text_encoder, dataloader, optimizer, device): model.train() total_loss 0.0 for batch in dataloader: images batch[image].to(device) masks batch[mask].to(device) text_feat batch[text_feat].to(device) optimizer.zero_grad() # 1. 图像编码 image_emb model.image_encoder(images) # 2. 调用SAM的prompt encoder生成稀疏/稠密提示 sparse_emb, dense_emb model.prompt_encoder( pointsbatch[point_coords], labelsbatch[point_labels], boxesbatch[boxes], masksNone ) # 3. 文本特征投影并与稀疏提示拼接 text_emb model.text_proj(text_feat).unsqueeze(1) # [B, 1, 256] sparse_emb torch.cat([sparse_emb, text_emb], dim1) # 4. mask decoder预测 low_res_logits, _ model.mask_decoder( image_embeddingsimage_emb, image_pemodel.prompt_encoder.get_dense_pe(), sparse_prompt_embeddingssparse_emb, dense_prompt_embeddingsdense_emb, multimask_outputFalse ) # 5. 上采样到与原图mask一致的分辨率 logits F.interpolate( low_res_logits, size(masks.shape[-2], masks.shape[-1]), modebilinear, align_cornersFalse ) loss focal_loss(logits, masks) dice_loss(logits, masks) loss.backward() optimizer.step() total_loss loss.item() return total_loss / len(dataloader)这里有一个坑需要提醒SAM的prompt_encoder中对点坐标有归一化和额外编码步骤不同版本的segment-anything库接口不太一样。如果你用的是官方最新版points参数格式是(B, N, 2)的tensorlabels是(B, N)。如果传入的坐标是numpy数组需要先转成tensor并指定dtype为float32。训练参数我设置了AdamW优化器初始学习率1e-4权重衰减1e-2batch size为8一共训练20个epoch。学习率采用余弦退火策略。在验证集上观察Dice指标每个epoch结束后保存最优权重。4.4 推理脚本与后处理流程推理脚本的流程比训练简单很多但后处理这块做不好效果会大打折扣。完整的推理流程是读图、提取图像特征、构造提示文本或点框、过mask decoder得到logits、sigmoid转概率、二值化、连通域处理。为什么需要连通域处理因为SAM在小目标或者低对比度场景下偶尔会输出一些零散的噪点区域这些区域面积很小和真正目标不相连。用OpenCV的connectedComponentsWithStats统计每个连通域的面积只保留面积大于阈值的区域就能把这些噪点滤掉。推理端我封装成了一个命令行调用python infer.py \ --image demo/0001.jpg \ --text 广告牌 \ --checkpoint checkpoints/sam_best.pth \ --output output/0001_mask.png如果走点框提示则传--point和--box参数。推理时默认用vit_b骨架如果加载的是已经训练好的权重直接替换checkpoint路径就能用。5. 已训练模型实测与部署5.1 模型权重与测试环境说明这个项目附带的权重是20个epoch训练出来的训练数据是广告牌和无类别目标混合的数据集总共约5000张图像。硬件是单张RTX 3090显存24GB冻结image encoder的情况下每个epoch大约15分钟整个训练一天多一点跑完。实测在400张验证图上的指标如下指标数值Dice0.911mIoU0.862F1-score0.905这几个数字不是拍脑袋写的是验证集上的真实结果。当然不同数据集、不同目标类别上指标会有波动如果你的数据分布差异很大建议重新微调而不是直接硬套。5.2 快速复现的完整步骤项目拿到手之后按下面四步操作就能复现训练和推理第一步下载代码和数据安装依赖库。数据目录里有标注好的json和对应的图像、mask文件权重文件放在checkpoints目录。第二步运行训练脚本python src/train.py --config config/train_config.yaml第三步用验证集评估模型效果评估脚本会输出Dice、mIoU和F1指标。第四步用推理脚本跑单张图。这一步我重点验证了文本提示的可控性同一张街道图文本分别换成广告牌、汽车、行人输出的掩膜确实会是图中的不同目标。这说明文本分支学到的不只是简单记忆而是真正建立起了语义描述和图像区域之间的关联。5.3 实测效果与当前局限实测效果好的场景有广告牌这类刚性强、纹理清晰的目标分割结果非常干净边缘贴合度高车辆、行人这类常见目标效果也不错。但有两个场景表现一般一是有遮挡的目标比如前景把广告牌挡住了一半SAM本身在这类场景下就会退化因为它的训练数据里遮挡样本比例不高二是文本描述和图像语义存在二义性时比如红色的图标如果图里同时有多个红色图标模型只能依赖文本的模糊语义来做判断结果不稳定。解决的办法也有在双模态基础上再加一个点/框提示来消除歧义。这也是我把点框和文本都做进同一个模型的原因它们是互补的不是替代关系。6. 常见问题与排查技巧实录我直接把实测过程中遇到的高频问题整理成一个速查表方便你以后遇到类似问题时快速定位。问题现象可能原因解决方案训练时显存OOMimage encoder参数量大或batch size过大冻结image encoder使用混合精度batch size降到4或2训练loss不下降学习率过大、样本中mask过小降低学习率到5e-5过滤掉mask面积小于1%图像面积的样本mask边缘粗糙不贴合上采样后没有细化或dice loss权重不够提高dice loss权重到0.8增加一个CRF后处理可选文本提示和图像区域对不上CLIP文本特征没有参与训练或投影层随机初始化先单独训练文本投影层10个epoch再联合微调decoder推理时报维度错误prompt encoder的输入格式不对检查points和labels是否为tensor坐标是否归一化到[0,1]训练时验证集Dice高但推理效果差数据增强过强导致推理时分布偏移推理时保持原始分辨率不做随机缩放和裁剪其中有两个问题我想展开细说因为它们最容易让人踩坑。第一个是文本投影层训练不充分的问题。如果一开始就让文本投影层和mask decoder同时训练投影层可能会在训练初期产生随机噪声导致模型根本学不到文本和图像区域对齐这个关系。我的经验是分两个阶段第一阶段冻结其他所有模块单独训练文本投影层让CLIP特征先适应SAM的256维空间第二阶段再一起微调projection和mask decoder。这样训练稳定很多Dice指标能高出两三个点。第二个是mask过小导致的loss不稳定问题。训练数据里如果有很多目标的mask只占整张图的不到0.5%dice loss在数值上会非常敏感一个像素的偏移就能让loss剧烈波动。处理方式是数据预处理时加一道过滤mask面积占比小于阈值的样本要么扔掉要么在采样时做re-sampling提高大目标的比例。这类小目标样本不是不能学而是需要单独调参不要混在一起训练。还有一个经验是如果需要在新的数据域上继续训练加载官方SAM权重做初始化会比加载我这个项目的权重更快收敛。官方权重是在通用数据上预训练的领域偏置小我训练过的权重偏向广告牌数据域在类似场景上效果好但换到医疗或者遥感场景时反而可能不如通用权重。跑实验的时候一定要区分清楚你是要跨领域泛化还是要在当前领域追求极致性能。另外一个实操细节在保存checkpoint时不要只保存model.state_dict()还要把模型的配置参数一起保存比如用了vit_b还是vit_l、图像输入尺寸、是否包含文本分支等。否则换个环境加载权重时模型结构对不上会非常痛苦。我在项目里是保存了一个dict包含model_state_dict、config和epoch恢复起来非常省事torch.save({ model_state_dict: model.state_dict(), config: config, epoch: epoch, best_dice: best_dice, }, save_path)最后再分享一个小技巧。如果你只是想在某个具体场景里快速验证效果不需要自己训练直接用项目里的推理脚本跑通一次然后用OpenCV做一个简单的批处理脚本把整个文件夹的图片都跑一遍结果输出到指定目录。很多工业场景里第一步要的不是一个漂亮的模型而是一个能证明这事可行的demo。跑通demo之后再决定是否需要采集更多数据来微调这个节奏最务实。本文还有配套的精品资源点击获取
返回列表