SAM 分割一切模型
SAM(Segment Anything Model)是 Meta 提出的通用图像分割基础模型,通过提示驱动(prompt-driven)的分割范式,实现了零样本迁移到新任务和新图像域,被誉为"分割领域的 GPT 时刻"。
架构设计
SAM 由三个核心组件构成:
- 图像编码器:基于 ViT-H 的视觉骨干,提取图像特征图
- 提示编码器:将点、框、掩码和文本等提示编码为嵌入向量
- 掩码解码器:轻量级 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)