
Segment Anything 图像分割模型前向传播全链路拆解【免费下载链接】segment-anythingThe repository provides code for running inference with the SegmentAnything Model (SAM), links for downloading the trained model checkpoints, and example notebooks that show how to use the model.项目地址: https://gitcode.com/GitHub_Trending/se/segment-anything在一张皮卡的照片上点了一个像素坐标Segment AnythingSAM从点、框提示生成物体掩码的图像分割模型内部接下来发生什么答案是图像特征早在点击前就算好并存着点击只做了一次位置编码和几层轻量 Transformer。整个模型由图像编码器、提示编码器、掩码解码器三部分组成串联逻辑在 sam.py 的Sam.forward里。下面按一次前向传播的顺序把链路拆开。它要替代什么传统语义分割和实例分割的瓶颈换一个任务就要重新标注一批像素级数据、重新训练一个模型流程从头再来一遍。Segment Anything 的取舍是把看懂图和听指令拆开——训练一个通用图像编码器把图像特征算成可复用的表示分割本身退化成一个轻量的提示到掩码映射。一张图的特征只算一次之后随便点、随便换提示。一次前向传播的全链路整条链路分五步预处理 → 特征提取 → 提示注入 → 掩码预测 → 上采样输出。前两步只跟图像有关后三步只跟提示有关这正是它支持交互的前提。输入图像归一化与方形补齐这一步输入是原始 HWC 图像输出是 1024×1024 的归一化张量。ResizeLongestSide先把长边缩到 1024Sam.preprocess再做两件事减 ImageNet 均值、除标准差右下角补零成正方形。def preprocess(self, x: torch.Tensor) - torch.Tensor: x (x - self.pixel_mean) / self.pixel_std # 归一化 h, w x.shape[-2:] padh self.image_encoder.img_size - h # 补到 1024×1024 padw self.image_encoder.img_size - w x F.pad(x, (0, padw, 0, padh)) return x补齐时把原始尺寸单独存下来输出阶段靠它裁掉填充区域保证掩码和原图像素严格对齐。图像特征编码ViT 主干加颈部网络这一步输入是 1024×1024 张量输出是 1×256×64×64 的图像特征图。ViT-H 配置下16×16 的 patch 卷积把图像切成 4096 个 token加可学习绝对位置嵌入过 32 层 Transformer 块最后经颈部网络Neck把高维特征压到低维的过渡层降维。def forward(self, x: torch.Tensor) - torch.Tensor: x self.patch_embed(x) # 16x16 卷积切 patch4096 个 token if self.pos_embed is not None: x x self.pos_embed for blk in self.blocks: x blk(x) # 32 层窗口/全局混合注意力 x self.neck(x.permute(0, 3, 1, 2)) # 1280 - 256 通道 return x这里有个细节容易忽略——32 层里 28 层只做 14×14 的窗口注意力只有第 8、16、24、32 层是全局注意力。窗口内计算量与窗口大小平方成正比而不是与 token 总数平方成正比这是 1024 输入下 ViT 能跑得动的关键。用户提示注入点、框、掩码统一成嵌入这一步输入是提示点坐标加标签、XYXY 框、或上轮掩码和上一步的特征图输出两类嵌入稀疏的点/框 token 序列 B×N×256稠密的掩码特征图 B×256×64×64。点先平移 0.5 到像素中心再经随机位置编码然后按标签加上可学习嵌入0 是负点、1 是正点、-1 是填充点。def _embed_points(self, points, labels, pad: bool) - torch.Tensor: points points 0.5 # 平移到像素中心 point_embedding self.pe_layer.forward_with_coords( points, self.input_image_size) # 随机频率正弦余弦编码 point_embedding[labels -1] self.not_a_point_embed.weight point_embedding[labels 0] self.point_embeddings[0].weight # 负点 point_embedding[labels 1] self.point_embeddings[1].weight # 正点 return point_embedding框的两个对角点走同一套位置编码分别加第 3、4 个可学习嵌入掩码输入走 3 层卷积下采样到 64×64没有掩码时用可学习的no_mask_embed铺满整张特征图。位置由位置编码承担、前景还是背景由可学习嵌入承担两者解耦后同一种编码方式就能通吃点和框。掩码预测双向 Transformer 加线性探针这一步输入是图像特征、稠密位置编码和两类提示嵌入输出 32×32 的掩码 logit 图和每张掩码的质量分数。MaskDecoder把 1 个 IoU token、4 个掩码 token 与提示 token 拼成查询序列图像特征加稠密提示嵌入当 key过一个 2 层的双向 Transformer——查询看图像图像也反过来看查询。output_tokens torch.cat([self.iou_token.weight, self.mask_tokens.weight], dim0) tokens torch.cat((output_tokens, sparse_prompt_embeddings), dim1) src torch.repeat_interleave(image_embeddings, tokens.shape[0], dim0) src src dense_prompt_embeddings # 提示条件注入图像特征 hs, src self.transformer(src, pos_src, tokens) # 2 层双向 Transformer # 每个掩码 token 经超网络生成一组权重与上采样特征图做内积 masks (hyper_in upscaled_embedding.view(b, c, h * w)).view(b, -1, h, w) iou_pred self.iou_prediction_head(iou_token_out)换句话说掩码不是卷积头吐出来的而是token 生成的权重向量与图像特征图的线性组合4 个掩码 token 相当于 4 个可学习线性探针同时输出 4 张候选掩码IoU 头再逐张打分。上采样输出裁掉填充、对齐原图这一步输入是 32×32 掩码 logit 加输入/原始尺寸输出是原图尺寸的布尔掩码、质量分数和 256 分辨率 logit。流程是先双线性插值到 1024×1024裁掉补齐区域再插值回原始尺寸用阈值 0.0 转布尔。masks F.interpolate(masks, (1024, 1024), modebilinear, align_cornersFalse) masks masks[..., : input_size[0], : input_size[1]] # 裁掉 padding masks F.interpolate(masks, original_size, modebilinear, align_cornersFalse)解码器只出低分辨率 logit重采样和裁剪全放后处理解码器因此保持轻量同时 256 分辨率的 logit 会原样返回供下一轮预测当mask_input回喂。三个值得聊的设计决策混合注意力为什么这样配选了什么ViT-H 的 32 层里 28 层用 14×14 窗口注意力4 层第 8、16、24、32用全局注意力。为什么选它窗口把注意力的计算量从与 token 数平方成正比降到与窗口大小平方成正比1024 输入、4096 token 下 ViT-H 才跑得动。代价是什么窗口内看不到窗口外跨区上下文全靠那 4 个全局层传递删掉任何一层精度都会掉。随机位置编码怎么工作选了什么不用可学习的网格表用一个固定的高斯随机矩阵做频率基坐标投影后取正弦余弦。为什么选它任意坐标都能直接算出编码与分辨率无关提示坐标和 64×64 特征图共用同一套编码函数。代价是什么不针对特定分辨率优化位置信息表达上限不如专门训练的表格换来的是提示坐标、图像特征、缩放填充全部不用查表。多掩码加 IoU 评分选了什么解码器固定 1 个 IoU token 加 4 个掩码 token其中第 1 张是模型自认最佳后 3 张是候选IoU 头逐张预测质量分数。为什么选它提示含糊时比如单点点在物体中间模型可以输出大中小三档轮廓由用户或下游逻辑按分数挑含糊性被显式建模。代价是什么token 数固定提示已经很明确时也要白算 4 张掩码好在解码器只有 2 层双向 Transformer这部分开销不大。最小可运行示例核心调用 5 行需要本地有 checkpoint 文件vit_h/l/b 三档from segment_anything import SamPredictor, sam_model_registry sam sam_model_registryvit_b predictor SamPredictor(sam) predictor.set_image(image_rgb_np) masks, iou, low_res predictor.predict(point_coords[[512, 384]], point_labels[1])predict返回三个值原图尺寸的二值掩码、对应的 IoU 质量分数、256×256 的 logit可直接回喂下一轮。完整演示在 predictor_example.ipynb点选交互和 automatic_mask_generator_example.ipynb整图批量生成ONNX 导出脚本是 export_onnx_model.py。往下游走这套架构对不需要新标注的任务直接可用自动掩码生成器批量产出全图物体掩码图像编辑、目标检测都能直接消费这些掩码。一个值得自己试的方向是把low_res作为mask_input回喂predictor.predict观察第二轮掩码如何收敛。命令行入口在 scripts/amg.pypython scripts/amg.py --checkpoint ckpt --model-type vit_b --input img --output dir。【免费下载链接】segment-anythingThe repository provides code for running inference with the SegmentAnything Model (SAM), links for downloading the trained model checkpoints, and example notebooks that show how to use the model.项目地址: https://gitcode.com/GitHub_Trending/se/segment-anything创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考