1. SAM模型架构解析与模块定位
在计算机视觉领域,Segment Anything Model(SAM)作为Meta推出的通用图像分割模型,其工程实现值得深入研读。ultralytics团队提供的PyTorch实现版本中,核心逻辑分布在amg.py、build.py、build_sam3.py和model.py四个子模块,每个文件各司其职又紧密协作。作为长期从事CV模型开发的工程师,我认为理解这些模块的协作关系是掌握SAM二次开发的关键前提。
从功能划分来看:
- amg.py:实现Automatic Mask Generation(自动掩码生成)流程,包含点网格生成、掩码后处理等非神经网络部分
- build.py:模型构建的入口文件,提供统一的模型工厂函数
- build_sam3.py:SAM-v3版本的特殊构建逻辑(与默认版本存在架构差异)
- model.py:包含图像编码器、提示编码器、掩码解码器等核心网络结构
提示:阅读源码时建议按照build.py→model.py→amg.py的顺序,先理清模型构造流程再研究具体实现。build_sam3.py仅在需要使用SAM-v3时重点查看。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 核心模块代码深度解析
2.1 build.py的工厂模式实现
作为模型构造的入口点,build.py通过build_sam函数封装了模型初始化过程。其核心逻辑如下:
python复制def build_sam(
checkpoint: Optional[str] = None,
model_type: str = "vit_h",
...
) -> Sam:
from .model import Sam # 延迟导入避免循环依赖
# 根据模型类型设置参数
if model_type == "vit_h":
encoder_embed_dim = 1280
encoder_num_heads = 16
elif model_type == "vit_l":
...
# 初始化SAM实例
sam = Sam(
image_encoder=ImageEncoderViT(...),
prompt_encoder=PromptEncoder(...),
mask_decoder=MaskDecoder(...),
...
)
# 加载预训练权重
if checkpoint is not None:
with open(checkpoint, "rb") as f:
state_dict = torch.load(f)
sam.load_state_dict(state_dict)
return sam
关键设计要点:
- 延迟导入机制:在函数内部导入Sam类,避免模块循环依赖
- 参数化配置:根据model_type动态设置ViT的embed_dim等超参数
- 权重加载分离:构造空模型后再加载预训练参数,保持灵活性
避坑指南:实际部署时需要注意PyTorch的版本兼容性问题。当从不同框架转换权重时,建议先用原始代码加载保存为PyTorch原生格式。
2.2 model.py的三阶段架构
model.py实现了SAM的核心网络结构,主要包含三个组件:
2.2.1 ImageEncoderViT(图像编码器)
基于Vision Transformer的改进实现,其特殊之处在于:
- 使用
kernel_size=16的卷积进行patch嵌入(而非常规的线性投影) - 采用
window_attention机制降低计算复杂度 - 输出多尺度特征图供解码器使用
python复制class ImageEncoderViT(nn.Module):
def __init__(
