Meta 的 Segment Anything Model (SAM) 是一个基于提示的通用图像分割模型,能够根据点、框、文本等交互式提示,对图像中的任何对象生成高质量的分割掩码。其核心在于零样本迁移能力,无需针对特定任务进行训练即可在新领域(如遥感图像)直接应用。

一、模型架构与工作原理SAM 的架构主要由三个组件构成:

  1. 图像编码器:基于 Vision Transformer (ViT),将输入图像编码为高维特征向量。
  2. 提示编码器:将用户提供的交互式提示(如点、框、文本)编码为提示嵌入。
  3. 掩码解码器:结合图像特征和提示嵌入,实时生成高质量的对象分割掩码。

其工作流程是:图像编码器一次性处理整张图像生成图像嵌入;当用户提供提示后,提示编码器将提示转换为向量;掩码解码器接收这两种嵌入,快速(在50ms内)输出多个可能的分割掩码及其置信度。

二、模型系列与演进

SAM 系列模型持续迭代,功能不断增强。

版本核心特性与升级主要应用场景
SAM (v1)奠定提示分割基础,支持点、框等空间提示。通用图像交互式分割。
SAM 2引入视频分割能力,优化模型架构与长上下文处理。图像与视频对象分割。
SAM 3原生支持自然语言文本提示,实现“文本到掩码”的跨模态分割。更直观的文本引导图像与视频分割。

此外,社区还推出了多种改进版本以适应不同需求:

  • FastSAM: 使用 CNN 替代 Transformer,大幅提升推理速度。
  • MobileSAM: 通过知识蒸馏得到的小型化模型,便于移动端部署。
  • EfficientSAM: 在效率和精度间取得平衡。
  • HQ-SAM: 专注于提升分割掩码的边界质量。

三、快速使用与部署

您可以通过多种方式快速体验 SAM,尤其是最新的 SAM 3。

1. 使用在线 Web 演示 (SAM 3)
访问 CSDN 星图平台 或 Meta 官方演示,通常提供一键部署的 Gradio 或 Streamlit Web 界面。您只需上传图像,并输入英文文本提示(如 “a white car”),即可获得分割结果。

2. 通过代码调用 (以 SAM v1 为例)
以下是使用官方 segment-anything Python 库进行交互式分割的基本步骤。

# 安装必要库
# pip install git+https://github.com/facebookresearch/segment-anything.git
# pip install opencv-python pycocotools matplotlib onnxruntime onnx

import numpy as np
import cv2
import matplotlib.pyplot as plt
from segment_anything import SamPredictor, sam_model_registry

# 1. 加载模型(以 ViT-H 大模型为例,需提前下载权重文件 sam_vit_h_4b8939.pth)
model_type = "vit_h"
checkpoint_path = "./sam_vit_h_4b8939.pth"
sam = sam_model_registry[model_type](checkpoint=checkpoint_path)
predictor = SamPredictor(sam)

# 2. 准备图像并生成图像嵌入
image = cv2.imread('your_image.jpg')
image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)
predictor.set_image(image) # 核心:计算并存储图像嵌入

# 3. 定义提示(例如:一个前景点坐标,格式为[x, y])
input_point = np.array([[500, 375]]) # 假设这是图像中某个对象上的点
input_label = np.array([1]) # 1 表示前景点,0 表示背景点

# 4. 预测分割掩码
masks, scores, logits = predictor.predict(
    point_coords=input_point,
    point_labels=input_label,
    multimask_output=True, # 输出多个可能的掩码)

# 5. 可视化最佳掩码
best_mask_idx = np.argmax(scores)
best_mask = masks[best_mask_idx]
plt.imshow(image)
plt.imshow(best_mask, alpha=0.5) # 半透明叠加
plt.scatter(input_point[:,0], input_point[:, 1], c='red', s=50) # 标出提示点
plt.show()

四、高级应用:微调与特定领域适配

当基础模型在特定领域(如医学影像、遥感解译)表现不佳时,可以进行微调。

微调核心步骤:

  1. 准备数据:将自定义数据集的标注(掩码)转换为模型可接受的提示格式(如从掩码中自动生成边界框或关键点作为提示)。
  2. 加载预训练模型:使用 SAM 的官方架构和预训练权重作为起点。
  3. 定义训练循环:冻结图像编码器(以保留通用视觉知识),主要训练提示编码器和掩码解码器。
# 微调伪代码示例(基于 PyTorch)
import torch
from torch.utils.data import DataLoader
from segment_anything import sam_model_registry

# 1. 加载预训练模型
model_type = "vit_h"
sam = sam_model_registry[model_type](checkpoint="./sam_vit_h_4b8939.pth")
sam.train() # 切换到训练模式

# 2.冻结图像编码器,只训练提示编码器和掩码解码器
for name, param in sam.image_encoder.named_parameters():
    param.requires_grad = False

# 3. 准备自定义数据集加载器(需实现 __getitem__ 返回 image, prompt, gt_mask)
dataset = YourCustomDataset(root='path/to/your/data')
dataloader = DataLoader(dataset, batch_size=4, shuffle=True)

# 4. 定义优化器和损失函数
optimizer = torch.optim.Adam(sam.prompt_encoder.parameters(), lr=1e-4)
criterion = torch.nn.BCEWithLogitsLoss() # 二值掩码常用损失

# 5. 训练循环
for epoch in range(num_epochs):
    for images, prompts, gt_masks in dataloader:
        # 前向传播 masks, scores, logits = sam(images, prompts, multimask_output=False)
        # 计算损失(假设 logits 是模型原始输出)
        loss = criterion(logits, gt_masks.float())
        # 反向传播与优化
        optimizer.zero_grad()
        loss.backward()
        optimizer.step()

完成微调后,可将模型集成到标注工具(如 ISAT)中,实现高效的半自动标注流水线。

五、典型应用场景

  • 遥感图像分析:快速提取建筑物、农田、水体等地理要素,掩码可转换为 GIS 矢量格式。
  • 图像编辑与创作:实现智能抠图、背景替换。
  • 视频对象分割与跟踪:结合 SAM 2 的视频能力,对视频序列中的目标进行分割与追踪。
  • 医疗影像分析:辅助分割细胞、器官或病灶区域。
  • 自动驾驶:用于道路场景中动态物体的识别与分割。

参考来源

 

Logo

DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。

更多推荐