☰
Mobile SAM轻量化部署:从蒸馏原理到ONNX端侧推理实战
2026/9/29 18:32:33 网站建设 项目流程

简介:这份资源是面向计算机视觉开发者与AI应用实践者的Mobile SAM模型文件包,针对在AnyLabeling标注工具中集成轻量级分割能力的需求,解决移动端或低算力环境下运行Segment Anything模型的问题。压缩包共3个文件,包含2个onnx格式的模型权重与1个yaml配置文件,整体约34.96MB,体积小巧便于快速部署。onnx文件分别承担图像编码与掩码解码任务,yaml则用于声明模型结构与参数,解压至指定模型目录后即可被AnyLabeling直接调用。目前已有484人学习下载,适合希望低成本体验自动分割标注、研究轻量化SAM推理流程的读者参考,可帮助理解编码器与解码器分离的部署思路,并快速搭建可用的交互式分割环境。

1. 拆开 mobile-sam-20230629.zip:一个能在笔记本上跑的 SegmentAnything

如果你之前试过在本地跑 SegmentAnything,大概率经历过显存不够、推理慢到怀疑人生的阶段。原版 SAM 的 ViT-H 权重接近 2.5GB,单张图在普通显卡上动辄几秒,想集成到自己的标注工具或边缘设备里基本不现实。而 anylabling 放出的 mobile-sam-20230629.zip,本质是把 Mobile SAM 这套轻量化方案连同权重、推理脚本、ONNX 导出一起打包,让你在一台没有独显的笔记本上也能把「点一下就把目标抠出来」这件事跑通。它解决的不是精度天花板问题,而是「能不能在消费级硬件上实时用起来」的问题。适合三类人:想给标注流水线加自动预分割的算法工程师、要在端侧做交互式抠图的嵌入式开发者、以及想读一份干净轻量代码来理解 SAM 蒸馏思路的学生。下面我按自己实际部署的顺序,把这个包从环境到导出讲透。

2. Mobile SAM 凭什么把 ViT-H 塞进小模型:蒸馏路径与选型账

2.1 从 ViT-H 到 TinyViT 的蒸馏逻辑

原版 SAM 的结构是「图像编码器 + 提示编码器 + 掩码解码器」。真正吃算力的是图像编码器,ViT-H 有 6.32 亿参数,每张图只跑一次但代价极高。Mobile SAM 的做法是冻结提示编码器和掩码解码器,只把图像编码器换成一个 5.78M 参数的 TinyViT,然后让 TinyViT 去模仿 ViT-H 输出的图像嵌入。这里的关键是:蒸馏发生在嵌入层,而不是最终掩码层。也就是说 TinyViT 学的不是「怎么分割」,而是「怎么像 ViT-H 一样理解图像」,分割能力由后面共享的解码器提供。

这个选择带来两个直接后果。第一,训练成本低,因为解码器不用重训,蒸馏目标只是嵌入对齐,常见做法是 L1 加余弦相似度的组合损失。第二,推理时图像编码器从 ViT-H 的几十 GFLOPs 降到约 1.5GFLOPs 量级,单张图编码从几百毫秒降到十几毫秒。代价是精度会掉一点,在边缘、细小物体上尤其明显,但交互式分割场景里用户点一下就能修正,这个精度损失是可接受的。

2.2 为什么选 Mobile SAM 而不是别的轻量分割

市面上轻量分割方案不少,比如 FastSAM 走的是 YOLOv8-seg 全实例分割路线,一次出所有掩码,但它不支持任意点提示,交互性差。Mobile SAM 保留了 SAM 的提示范式,你可以给点、给框、给粗掩码,这对标注工具是刚需。另一个对比对象是直接量化原版 SAM,INT8 量化后 ViT-H 依然有几百 M 参数,移动端内存吃紧。所以如果你的需求是「交互式、可提示、端侧可跑」,Mobile SAM 是当前性价比最高的选择;如果你要的是「一次性全图分割、不需要交互」,那 FastSAM 更合适。选型时先问自己:用户会不会给提示?会,就 Mobile SAM。

2.3 包内结构与依赖版本核对

拿到 mobile-sam-20230629.zip 后先别急着解压跑,先看结构。常见布局是:mobile_sam/放模型定义,weights/放mobile_sam.pt,根目录有setup.py、README、以及导出 ONNX 的脚本。权重文件大约 40MB 左右,这是它最讨喜的地方。

unzip mobile-sam-20230629.zip -d mobile-sam cd mobile-sam find . -maxdepth 2 -name "*.py" -o -name "*.pt" | sort

这段命令先解压再列出关键文件,目的是确认权重和模型定义都在。参数上-d指定解压目录,-maxdepth 2避免翻进深层缓存目录。如果发现没有.pt文件,说明压缩包只含代码,需要按 README 单独下载权重,别硬跑。

依赖方面,核心是torch、torchvision、timm、numpy、opencv-python。TinyViT 的实现依赖 timm 的某些层,版本不匹配会报cannot import name 'window_partition'之类的错。我一般先建独立环境再装:

conda create -n mobilesam python=3.10 -y conda activate mobilesam pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install timm opencv-python numpy matplotlib onnx onnxruntime

CUDA 版本按自己显卡改,没有独显就用 CPU 版 torch。timm 不要锁太老的版本,0.9 以上基本兼容。装完先python -c "import timm, torch; print(timm.__version__, torch.__version__)"验证,别等到跑推理才暴露。

3. 用 mobile_sam 在本地跑通第一张图的自动掩码

3.1 加载模型与生成器的最小脚本

跑通第一张图是建立信心的关键。Mobile SAM 的 API 和原版 SAM 几乎一致,sam_model_registry里注册的键名通常是vit_t。下面是最小可运行脚本:

import cv2 import numpy as np import torch from mobile_sam import sam_model_registry, SamAutomaticMaskGenerator # 选 vit_t 这个轻量编码器,权重路径按实际解压位置改 sam = sam_model_registry["vit_t"](checkpoint="weights/mobile_sam.pt") device = "cuda" if torch.cuda.is_available() else "cpu" sam.to(device=device) sam.eval() # 自动掩码生成器,points_per_side 控制网格密度 mask_generator = SamAutomaticMaskGenerator( model=sam, points_per_side=16, # 默认 32,端侧降到 16 省算力 pred_iou_thresh=0.8, # 掩码质量阈值,越高越少但越准 stability_score_thresh=0.9, min_mask_region_area=100 # 过滤碎小区域 ) image = cv2.imread("test.jpg") image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) masks = mask_generator.generate(image) print("生成掩码数量:", len(masks))

逻辑上,sam_model_registry["vit_t"]实例化 TinyViT 版本并加载权重,SamAutomaticMaskGenerator会在图上铺网格点,每个点作为提示喂给解码器,再过滤重叠和低质量掩码。参数说明:points_per_side是网格边长,16 表示 256 个提示点,32 是 1024 个,端侧建议 12 到 16;pred_iou_thresh低于 0.8 会混入大量噪声掩码;min_mask_region_area对标注场景很有用,能砍掉几十像素的碎片。

3.2 交互式点提示:给一个点抠一个目标

自动掩码适合批量预标注,但真正高频的是「用户点一下」。这时用SamPredictor:

from mobile_sam import SamPredictor predictor = SamPredictor(sam) predictor.set_image(image) # 图像编码只跑一次,后续提示复用 # 用户点击的坐标,格式是 (x, y) input_point = np.array([[320, 240]]) input_label = np.array([1]) # 1 表示前景点,0 表示背景点 masks, scores, logits = predictor.predict( point_coords=input_point, point_labels=input_label, multimask_output=True # 输出 3 个候选掩码供选择 ) best = masks[np.argmax(scores)]

关键点在set_image只做一次编码,之后每次点击只跑轻量解码器,这就是交互能实时的原因。multimask_output=True会返回三个粒度不同的掩码,通常分数最高的那个是整体,另外两个可能是局部,UI 上让用户点选体验更好。如果只要一个结果,设成 False 并配合return_logits做后续 refine。

3.3 参数怎么调:points_per_side 与阈值组合

很多人跑完发现掩码要么太碎要么漏目标,问题基本出在参数组合。下面这张表是我在不同场景下试出来的经验值:

场景points_per_sidepred_iou_threshstability_score_threshmin_mask_region_area
端侧实时预览120.750.8550
标注预分割240.820.90100
精细小物体320.880.9220
大场景概览160.800.88500

调参顺序建议:先定points_per_side,它直接决定算力和召回;再调pred_iou_thresh控制质量;最后用min_mask_region_area清碎片。注意stability_score_thresh调太高会把边缘柔和的目标(比如毛发、烟雾)整片滤掉,这类场景反而要降到 0.85 左右。

4. 导出 ONNX 与端侧部署:把 PyTorch 权重变成可移植文件

4.1 导出编码器和解码器两个 ONNX

端侧部署不能带 PyTorch,得转 ONNX。Mobile SAM 要分成两部分导出,因为编码器只跑一次、解码器每次提示都跑,分开才能复用编码结果。

import torch from mobile_sam import sam_model_registry sam = sam_model_registry["vit_t"](checkpoint="weights/mobile_sam.pt") sam.eval() # 导出图像编码器,输入固定 1024x1024 dummy_image = torch.randn(1, 3, 1024, 1024) torch.onnx.export( sam.image_encoder, dummy_image, "mobile_sam_encoder.onnx", input_names=["image"], output_names=["image_embedding"], opset_version=17, dynamic_axes={"image": {0: "batch"}} ) # 导出提示编码器 + 掩码解码器,输入是嵌入、点坐标、点标签 dummy_embedding = torch.randn(1, 256, 64, 64) dummy_points = torch.randn(1, 1, 2) dummy_labels = torch.randint(0, 2, (1, 1)) torch.onnx.export( sam.mask_decoder, (dummy_embedding, sam.prompt_encoder.get_dense_pe(), dummy_points, dummy_labels), "mobile_sam_decoder.onnx", input_names=["image_embedding", "image_pe", "point_coords", "point_labels"], output_names=["masks", "iou_predictions"], opset_version=17 )

导出时opset_version建议 17,低版本对某些算子支持不全。dynamic_axes只对 batch 维开放,图像尺寸最好固定成 1024,因为位置编码是固定长度的,动态尺寸会报错。解码器导出时注意prompt_encoder.get_dense_pe()的输出也要作为输入传进去,否则位置信息丢失,掩码会整体偏移。

4.2 ONNXRuntime 推理与前后处理对齐

导出后要用 ONNXRuntime 验证数值是否和 PyTorch 一致,别直接上端侧。

import onnxruntime as ort import numpy as np sess_enc = ort.InferenceSession("mobile_sam_encoder.onnx", providers=["CPUExecutionProvider"]) sess_dec = ort.InferenceSession("mobile_sam_decoder.onnx", providers=["CPUExecutionProvider"]) # 前处理:归一化 + resize 到 1024 img = cv2.resize(image, (1024, 1024)).astype(np.float32) / 255.0 mean = np.array([0.485, 0.456, 0.406], dtype=np.float32) std = np.array([0.229, 0.224, 0.225], dtype=np.float32) img = (img - mean) / std img = img.transpose(2, 0, 1)[None, ...] embedding = sess_enc.run(None, {"image": img})[0] # 点坐标要归一化到 1024 尺度 point = np.array([[[320 / orig_w * 1024, 240 / orig_h * 1024]]], dtype=np.float32) label = np.array([[1]], dtype=np.float32) masks, iou = sess_dec.run(None, { "image_embedding": embedding, "image_pe": dense_pe, "point_coords": point, "point_labels": label })

前后处理最容易翻车的是归一化参数和坐标缩放。Mobile SAM 用的是 ImageNet 的 mean/std,和原版一致;点坐标必须按原图到 1024 的比例缩放,忘了这步掩码会跑到错误位置。dense_pe可以从 PyTorch 侧导出一次存成 npy,端侧直接加载,省得每次算。

4.3 端侧内存与线程配置

ONNXRuntime 在移动端默认线程数可能过多,反而拖慢。常见做法是限制 intra-op 线程:

options = ort.SessionOptions() options.intra_op_num_threads = 2 options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL sess = ort.InferenceSession("mobile_sam_encoder.onnx", options, providers=["CPUExecutionProvider"])

编码器峰值内存大约 200MB 出头,解码器几十 MB,中端手机能扛住。如果还紧张,可以对编码器做 INT8 动态量化,精度掉 1 到 2 个点,但内存和延迟都能再降一截。量化后务必用同一张图对比掩码 IoU,低于 0.9 就别用了。

5. 避坑与排查:mobile-sam 部署里最容易翻车的五件事

5.1 现象:报KeyError: 'vit_t',模型注册不上

原因通常是mobile_sam包没被正确安装或导入路径不对,sam_model_registry里没有注册 TinyViT。解决:确认在解压目录下执行,或pip install -e .把包装进环境;再检查mobile_sam/build_sam.py里是否有_build_sam_vit_t并注册了vit_t。如果用的是别处下载的权重配这份代码,键名可能叫vit_tiny,以代码里的注册名为准。

5.2 现象:掩码整体偏移或只覆盖一角

原因几乎都是坐标没归一化或图像预处理不一致。自动掩码生成器内部会处理,但你自己写交互推理时,点坐标必须缩放到模型输入尺度,且图像要按同样的 resize 和 padding 处理。解决:把原图到 1024 的缩放比例算清楚,点坐标乘同一个比例;如果用 letterbox padding,还要加上偏移量。建议先跑官方 demo 图确认基线,再换自己的图。

5.3 现象:CPU 上单张图要好几秒

原因可能是points_per_side设太高,或者没复用图像嵌入。自动掩码模式下每个点都要跑一次解码器,32 的网格就是 1024 次解码。解决:端侧把points_per_side降到 12 到 16;交互模式务必用SamPredictor的set_image只编码一次。另外确认 torch 没在跑多线程争抢,torch.set_num_threads(4)限制一下。

5.4 现象:ONNX 输出和 PyTorch 对不上

原因常见于 opset 版本、动态轴设置、以及解码器输入顺序。解决:先用固定输入(batch=1、尺寸 1024)导出,关掉所有 dynamic_axes,逐层对比中间输出;确认image_pe作为显式输入传入而不是在解码器内部生成;opset 统一用 17。数值误差在 1e-3 以内算正常,超过就是结构问题。

5.5 现象:小物体和细长目标分割断裂

这是 TinyViT 容量决定的固有短板,不是 bug。原因在于蒸馏后的嵌入对高频细节保留不足。解决:交互场景下多给几个点,前景点加背景点组合能显著改善;自动模式下提高points_per_side并降低min_mask_region_area;如果业务强依赖细粒度,考虑对目标区域裁剪后再送模型,用局部放大换精度。

6. 把 Mobile SAM 接进标注流水线的进阶技巧

单张推理跑通只是起点,真正省时间的是把它接进标注工具做预分割。我的做法是:先用SamAutomaticMaskGenerator以较低points_per_side(比如 12)跑一遍全图,得到粗掩码作为候选;标注员在 UI 上点选候选框,再用SamPredictor做一次点提示精修。这样 90% 的常见目标一次点选就成型,只有边缘复杂的才手动修。实测在 1080p 图上,粗分割加精修的单图交互时间能压到 1 秒以内。

验证方案是否值得投入,我一般看两个指标:一是预分割掩码被标注员直接采纳的比例,低于 60% 说明参数或场景不匹配;二是端侧单次编码延迟,超过 200ms 交互就会卡顿。这两个数达标,这套方案就值得往流水线里塞。

还有一个容易被忽略的技巧:把dense_pe和常用分辨率的前处理参数缓存成文件,端侧启动时直接加载,省掉每次初始化的开销。我在一个嵌入式项目里就是这么干的,冷启动从 1.8 秒降到 0.6 秒。血泪经验是别在端侧动态生成位置编码,那玩意儿又慢又容易和训练时对不齐。

最后说个习惯:每次换权重或改导出脚本,我一定先用同一张测试图跑一遍,把掩码 IoU 和延迟记在表格里,和上一版对比。没有这个基线,你根本不知道某次「优化」到底是进步还是退步。Mobile SAM 这类轻量模型,参数和硬件的组合太多,靠感觉调就是玄学,靠数据调才踏实。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询