Skip to content

SAM 分割一切模型

SAM(Segment Anything Model)是 Meta 提出的通用图像分割基础模型,通过提示驱动(prompt-driven)的分割范式,实现了零样本迁移到新任务和新图像域,被誉为"分割领域的 GPT 时刻"。

SAM架构示意

架构设计

SAM 由三个核心组件构成:

  1. 图像编码器:基于 ViT-H 的视觉骨干,提取图像特征图
  2. 提示编码器:将点、框、掩码和文本等提示编码为嵌入向量
  3. 掩码解码器:轻量级 Transformer 解码器,融合图像和提示特征生成分割掩码
python
import torch
import torch.nn as nn

class PromptEncoder(nn.Module):
    """SAM 提示编码器"""
    def __init__(self, embed_dim=256):
        super().__init__()
        self.point_embeddings = nn.ModuleList([
            nn.Embedding(1, embed_dim) for _ in range(4)
        ])  # 正/负前景/背景点
        self.box_embedding = nn.Sequential(
            nn.Linear(4, embed_dim), nn.ReLU(), nn.Linear(embed_dim, embed_dim)
        )
        self.no_mask_embed = nn.Embedding(1, embed_dim)

    def forward(self, points=None, boxes=None, masks=None):
        prompts = []
        if points is not None:
            coords, labels = points
            point_embed = sum(
                self.point_embeddings[i] for i in labels
            )
            prompts.append(coords + point_embed)
        if boxes is not None:
            prompts.append(self.box_embedding(boxes))
        return torch.cat(prompts, dim=1)

class MaskDecoder(nn.Module):
    """SAM 掩码解码器"""
    def __init__(self, num_multimask_outputs=3):
        super().__init__()
        self.transformer = TwoWayTransformer()
        self.output_upscaling = nn.Sequential(
            nn.ConvTranspose2d(256, 128, 2, 2),
            nn.GELU(),
            nn.ConvTranspose2d(128, 64, 2, 2),
        )
        self.output_hypernetworks = nn.ModuleList([
            MLP(256, 256, 64*64) for _ in range(num_multimask_outputs + 1)
        ])

    def forward(self, image_embeddings, prompt_embeddings):
        # 双向 Transformer 交互
        hs = self.transformer(image_embeddings, prompt_embeddings)
        # 上采样并生成掩码
        upscaled = self.output_upscaling(hs)
        masks = [mlp(hs) for mlp in self.output_hypernetworks]
        return masks

提示类型

SAM 支持多种灵活的提示方式:

提示类型输入格式说明
前景点(x, y), label=1指定目标上的点
背景点(x, y), label=0指定非目标上的点
边界框(x1, y1, x2, y2)框选目标区域
掩码二值掩码粗糙的初始掩码
文本文本描述需配合其他模型

模糊感知输出

SAM 的掩码解码器默认输出 3 个候选掩码(单对象、包含对象、整体对象),以解决提示的歧义性。单点提示时,模型自动选择最合理的掩码。

SA-1B 数据集

SAM 的训练数据 SA-1B 是图像分割领域迄今最大的数据集:

  • 1100 万张高分辨率图像
  • 11 亿个分割掩码,平均每张图像约 100 个
  • 掩码由自动标注流程生成,经人工验证质量可靠

SAM 2:视频分割

SAM 2 将分割能力从静态图像扩展到视频:

  • 流式记忆架构:维护视频帧的记忆库,支持前后向传播
  • 实时视频分割:首帧标注后,后续帧自动跟踪分割
  • SAM 2.1:改进了小目标和遮挡场景的处理
python
from sam2 import SAM2ImagePredictor, SAM2VideoPredictor

# 图像分割
image_predictor = SAM2ImagePredictor.from_pretrained("facebook/sam2-hiera-large")
image_predictor.set_image(image)
masks = image_predictor.predict(point_coords=points, point_labels=labels)

# 视频分割
video_predictor = SAM2VideoPredictor.from_pretrained("facebook/sam2-hiera-large")
with video_predictor.state(video_frames) as state:
    video_predictor.add_new_points(state, frame_idx=0, points=points, labels=labels)
    video_masks = video_predictor.propagate_in_video(state)

相关资源

最近更新