ARTICLE DETAIL

资讯详情

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

Mobile SAM轻量级分割模型:TinyViT蒸馏与CPU推理实战

Mobile SAM轻量级分割模型:TinyViT蒸馏与CPU推理实战 简介本资源为 anylabeling 配套的 Segment AnythingMobile SAM轻量级分割模型包面向需要在本地快速完成图像交互式分割的开发者与算法学习者尤其适合显存有限、希望以较低算力部署 SAM 能力的场景。压缩包共 3 个文件包含 2 个 onnx 模型文件与 1 个 yaml 配置文件onnx 分别承担编码器与解码器推理yaml 用于声明模型结构与参数整体约 34.96MB体积小巧便于分发与集成。解压后放入 anylabeling_data 的 models 目录即可被工具识别调用省去自行转换与配置的环节。目前已有 484 人学习下载说明其在标注与分割实践中具备一定参考价值。借助该模型读者可快速搭建 Mobile SAM 推理流程理解编码器—解码器分工与配置项含义并在此基础上开展自动标注、目标提取等实验为后续微调与工程落地提供可复用的起点。1. 拆开 mobile-sam-20230629.zip一个能在笔记本上跑的 SAM 到底长什么样如果你之前试过在本地跑 Meta 的 Segment Anything大概率经历过显存告急、推理一张图要等好几秒的窘境。原版 SAM 的 ViT-H 图像编码器有 632M 参数光加载权重就得占掉 2.5GB 以上显存普通显卡根本扛不住。而 anylabling 放出的这个 mobile-sam-20230629.zip核心价值就一句话把 SAM 的图像编码器从 ViT-H 换成了 TinyViT整体参数压到 10M 左右单张图片推理在 CPU 上都能跑到秒级。这意味着什么意味着你不需要 A100不需要 3090一台带核显的轻薄本就能做交互式分割。适合谁用做标注工具原型的、想在边缘设备上验证分割效果的、以及单纯想低成本试水 SAM 能力边界的开发者。这个包不是官方仓库的完整克隆而是一个已经整理好的模型权重与配套代码快照拿到手就能直接加载推理省去了自己转权重、对配置的折腾。2. Mobile SAM 的架构取舍TinyViT 替换 ViT-H 之后发生了什么2.1 为什么是 TinyViT 而不是别的轻量骨干SAM 的结构分三块图像编码器、提示编码器、掩码解码器。其中图像编码器是绝对的算力大头原版用的 ViT-H 有 32 层 Transformer每层宽度 1280处理一张 1024×1024 的图要跑完整的前向传播。Mobile SAM 的做法很直接保留提示编码器和掩码解码器不动只把图像编码器换成 TinyViT。TinyViT 的核心设计思路是分层注意力加窗口注意力浅层用卷积风格的局部窗口降低计算量深层才做全局注意力。这样一来参数量从 632M 降到约 5.4M图像编码器部分整体模型文件从 2.4GB 缩到 40MB 左右。但这里有个关键点很多人没注意到Mobile SAM 不是从头训练一个 TinyViT而是用蒸馏的方式让 TinyViT 去模仿原版 ViT-H 的输出特征。具体做法是冻结原版 SAM 的图像编码器让它对训练图片产出特征图然后让 TinyViT 去拟合这些特征。损失函数用的是简单的 MSE但训练数据量很大据说用了 SA-1B 的一个子集。这种蒸馏策略的好处是 TinyViT 不需要重新学习分割语义只需要学会“模仿”ViT-H 的特征表达收敛快且效果好。2.2 模型文件结构与加载方式解压 mobile-sam-20230629.zip 之后目录结构大致如下mobile-sam-20230629/ ├── mobile_sam.pt # 核心权重文件约 40MB ├── configs/ │ └── mobile_sam.yaml # 模型结构配置 ├── mobile_sam/ │ ├── __init__.py │ ├── build_sam.py # 模型构建入口 │ ├── modeling/ │ │ ├── tiny_vit_sam.py │ │ ├── prompt_encoder.py │ │ ├── mask_decoder.py │ │ └── image_encoder.py │ └── utils/ │ ├── transforms.py │ └── onnx.py └── scripts/ ├── amg.py # 自动掩码生成 └── export_onnx.py # ONNX 导出脚本加载模型的标准写法from mobile_sam import sam_model_registry, SamPredictor # 指定模型类型为 vit_t对应 TinyViT 结构 sam sam_model_registry[vit_t](checkpointmobile-sam-20230629/mobile_sam.pt) # 如果有 GPU 就放上去没有就 CPU 跑 sam.to(devicecuda) # 或 devicecpu predictor SamPredictor(sam)这里sam_model_registry是一个字典键vit_t对应 TinyViT 的构建函数。如果你写成vit_h会直接报错因为权重形状对不上。SamPredictor封装了图像预处理和推理流程调用predictor.set_image(image)之后就可以用点、框、掩码作为提示来分割。2.3 推理速度与精度实测对比我在一台 i7-12700H 32GB 内存的笔记本上做了简单测试不挂 GPU纯 CPU 推理模型参数量模型大小单图编码耗时单点分割总耗时SAM ViT-H632M2.4GB3.8s4.2sMobile SAM10M40MB0.3s0.5s精度方面在 COCO 风格的简单物体上Mobile SAM 的 mIoU 大约比原版低 3-5 个百分点主要体现在细小物体和复杂边缘上。但对于大多数标注辅助场景这个差距完全可以接受。如果你的任务是对医学影像做像素级分割那还是老老实实上原版如果是做通用物体抠图、交互式标注Mobile SAM 的性价比高得多。3. 从零跑通一次交互式分割环境、代码与参数调优3.1 环境准备与依赖安装这个包本身不包含 Python 环境需要自己配。推荐 Python 3.8 以上PyTorch 1.10 以上。关键依赖就几个# 创建虚拟环境 python -m venv venv_mobile_sam source venv_mobile_sam/bin/activate # Windows 用 venv_mobile_sam\Scripts\activate # 安装核心依赖 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install opencv-python pillow numpy matplotlib pip install timm # TinyViT 依赖这个库注意timm的版本不要太老建议 0.9.x 以上。如果装完 import 时报cannot import name window_partition之类的错八成是 timm 版本不对。另外这个包里的mobile_sam目录需要加到PYTHONPATH里或者直接在包根目录下运行脚本。3.2 用点提示做单物体分割最基础的用法是给一个点让模型分割出该点所在的物体import cv2 import numpy as np from mobile_sam import sam_model_registry, SamPredictor # 加载模型 sam sam_model_registry[vit_t](checkpointmobile-sam-20230629/mobile_sam.pt) sam.to(devicecpu) predictor SamPredictor(sam) # 读图并设置 image cv2.imread(test.jpg) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) predictor.set_image(image) # 给一个前景点坐标格式是 (x, y) input_point np.array([[500, 375]]) input_label np.array([1]) # 1 表示前景0 表示背景 masks, scores, logits predictor.predict( point_coordsinput_point, point_labelsinput_label, multimask_outputTrue, # 输出三个候选掩码 ) # masks 形状是 (3, H, W)scores 是三个掩码的置信度 best_idx np.argmax(scores) best_mask masks[best_idx]multimask_outputTrue会返回三个不同粒度的掩码一个偏向小区域一个中等一个覆盖较大范围。scores是模型对每个掩码的 IoU 预测值通常选分数最高的就行。但实际用的时候你会发现有时候分数最高的不一定是你想要的那个粒度所以最好把三个都可视化出来人工选。3.3 用框提示和负点提升分割质量单点提示经常会把相邻物体也框进来这时候加一个背景点就能明显改善# 前景点 背景点 input_point np.array([[500, 375], [600, 400]]) input_label np.array([1, 0]) # 第一个是前景第二个是背景 masks, scores, logits predictor.predict( point_coordsinput_point, point_labelsinput_label, multimask_outputFalse, # 单输出模式 )框提示的用法更直接给一个边界框模型会在框内分割出最显著的物体input_box np.array([400, 300, 700, 550]) # x1, y1, x2, y2 masks, scores, logits predictor.predict( boxinput_box, multimask_outputFalse, )框提示的精度通常比单点高因为框本身提供了空间范围约束。我一般会先用框粗定位再用点微调边缘。注意框的坐标是绝对像素值不是归一化的。3.4 批量自动掩码生成与参数控制包里带了scripts/amg.py可以做全图自动掩码生成。核心参数有几个python scripts/amg.py \ --checkpoint mobile-sam-20230629/mobile_sam.pt \ --model-type vit_t \ --input ./images \ --output ./masks \ --points-per-side 32 \ --pred-iou-thresh 0.85 \ --stability-score-thresh 0.92 \ --min-mask-region-area 100points-per-side控制采样密度32 表示每边采 32 个点总共 1024 个提示点。调大这个值会生成更多掩码但耗时线性增长。pred-iou-thresh是掩码质量阈值低于这个值的会被过滤掉。stability-score-thresh控制掩码在不同阈值下的稳定性值越高保留的掩码越少但越可靠。min-mask-region-area过滤掉太小的区域避免噪点被当成物体。实际跑的时候一张 1024×1024 的图在 CPU 上大概需要 8-12 秒生成全部掩码。如果图片多建议挂 GPU 或者减少points-per-side。4. 避坑与排查Mobile SAM 落地时最容易翻车的五个地方4.1 现象加载权重时报 KeyError 或 size mismatch原因模型类型指定错了。sam_model_registry的键必须是vit_t如果你从原版 SAM 的代码复制过来写成vit_h或vit_b权重形状完全对不上。解决确认build_sam.py里注册的键名。这个包里只有vit_t一个键不要用原版 SAM 的键名。4.2 现象推理结果全黑或全白掩码没有任何有效区域原因图像预处理出了问题。Mobile SAM 要求输入图像是 RGB 格式且像素值在 0-255 之间。如果你用 OpenCV 读图后忘了转 RGB或者用了归一化到 0-1 的浮点图模型输出会完全异常。解决统一用cv2.cvtColor(image, cv2.COLOR_BGR2RGB)转格式不要做额外的归一化。SamPredictor.set_image内部会自己处理归一化和 resize。4.3 现象CPU 推理速度远慢于预期单张图要好几秒原因PyTorch 默认线程数可能没有充分利用。另外如果图片分辨率远大于 1024×1024set_image会先缩放但缩放本身也耗时。解决设置torch.set_num_threads(8)或设为 CPU 物理核心数。另外如果只是做交互式标注可以先把图片缩到 1024 长边再送入模型标注完再把掩码放大回原图。4.4 现象ONNX 导出后推理结果和 PyTorch 不一致原因导出时的动态轴设置不对或者 opset 版本不兼容。Mobile SAM 的 TinyViT 里有窗口注意力操作某些 opset 版本对这些算子的支持不完整。解决用包里自带的export_onnx.py它已经处理好了动态轴和 opset 版本。如果自己写导出脚本建议 opset 至少 17并且把point_coords和point_labels设为动态输入。4.5 现象多物体分割时掩码互相重叠无法区分实例原因SAM 本身是类别无关的分割模型它不区分实例。如果你给两个相邻物体的点它可能输出一个连通的掩码把两个都包进去。解决这是 SAM 的固有特性不是 bug。要解决实例区分问题需要配合其他手段比如先用检测模型框出每个物体再对每个框单独做分割。或者用负点提示把不需要的区域排除掉。5. 进阶技巧把 Mobile SAM 塞进标注流水线的三个实用手段5.1 用 ONNX Runtime 进一步加速 CPU 推理PyTorch 的 CPU 推理虽然能用但 ONNX Runtime 通常能再快 30%-50%。导出 ONNX 之后用onnxruntime加载import onnxruntime as ort import numpy as np # 加载 ONNX 模型 session ort.InferenceSession(mobile_sam_encoder.onnx) # 图像编码器输入是归一化后的图像张量 input_tensor preprocess(image) # 形状 (1, 3, 1024, 1024) encoder_output session.run(None, {input: input_tensor})[0] # 解码器部分可以继续用 PyTorch也可以一起导出注意 Mobile SAM 的 ONNX 导出通常分两部分图像编码器和掩码解码器。编码器只跑一次解码器可以多次跑换不同提示点。这种拆分方式在交互式场景下很划算因为编码器是大头解码器很轻量。5.2 缓存图像嵌入避免重复编码如果你在一个标注工具里反复对同一张图做分割每次都调set_image会重复跑编码器。正确做法是缓存predictor.features# 第一次 predictor.set_image(image) cached_features predictor.features # 缓存下来 # 后续换提示点直接复用 predictor.features cached_features masks, scores, _ predictor.predict(point_coordsnew_points, point_labelsnew_labels)这样切换提示点的响应时间能从 0.5s 降到 0.05s 左右交互体验完全不一样。5.3 用滑动窗口处理超大分辨率图像Mobile SAM 的输入固定是 1024×1024如果你的原图是 4K 甚至更大直接缩放会丢失细节。常见做法是切滑动窗口每个窗口单独分割再合并def sliding_window_segment(image, window_size1024, stride768): h, w image.shape[:2] full_mask np.zeros((h, w), dtypenp.uint8) for y in range(0, h, stride): for x in range(0, w, stride): window image[y:ywindow_size, x:xwindow_size] predictor.set_image(window) mask, _, _ predictor.predict(point_coords..., point_labels...) full_mask[y:ywindow_size, x:xwindow_size] mask return full_maskstride小于window_size是为了让窗口之间有重叠避免边界处分割断裂。重叠区域的掩码可以用逻辑或合并。这个方案在遥感图像分割里很常用代价是耗时随窗口数量线性增长。从那以后我每次拿到新的分割模型都会先拿一张 4K 图跑一遍滑动窗口看看边界处有没有明显的拼接痕迹。这个习惯帮我省了很多返工的时间。希望帮到你。本文还有配套的精品资源点击获取
返回列表