SAM模型实战:如何用自定义脚本实现更灵活的图像分割(Python代码详解)

如果你已经玩过Meta AI开源的Segment Anything Model(SAM),体验过官网Demo那种“点哪分哪”的炫酷,但很快发现,当你想把它集成到自己的项目流水线、处理一批特定格式的图片,或者需要更精细地控制分割后的输出时,官方的命令行工具和基础API就显得有些束手束脚了。你需要的不是另一个教程,告诉你如何运行amg.py,而是一把能够拆解、重组并强化SAM能力的“手术刀”。这正是本文要探讨的核心:抛开官方封装,我们如何通过编写自定义的Python脚本,深入SAM的预测核心,构建一个从交互逻辑、后处理到批量流水线都完全可控的图像分割工作流。这不仅仅是调用一个模型,而是真正将它“工程化”。

1. 超越官方Demo:为何需要自定义脚本?

官方的SAM仓库提供了强大的模型和基础的推理脚本,例如amg.py(自动掩码生成)和predictor.py提供的交互式预测接口。对于快速验证和简单应用,它们足够了。但当我们面对真实世界的项目时,几个痛点会立刻浮现:

  • 交互逻辑固化:官方示例的交互循环通常绑定在OpenCV的窗口事件里,难以嵌入到无界面的Web服务、自动化脚本或带有复杂UI的桌面应用中。
  • 输出结果单一:amg.py生成的是一堆二进制掩码文件,而实际应用可能需要带透明通道的PNG、裁剪后的对象图像、叠加了颜色标注的可视化图,或者直接的结构化数据(如轮廓坐标)。
  • 缺乏流程控制:批量处理时,我们可能希望对每张图片应用不同的提示策略(例如,对某类物体默认加一个负向点),或者在生成掩码后进行特定的过滤(如按面积、紧凑度筛选)。
  • 性能与调试:当处理成千上万张图片时,我们需要更细粒度的性能监控、错误处理以及中间结果的保存机制,以便回溯和优化。

自定义脚本的价值就在于,它让你从“API调用者”转变为“流程设计者”。你可以决定提示信息如何收集、模型推理前后进行何种处理、结果以何种形态交付。下面,我们将从一个基础但功能完整的交互式脚本出发,逐步拆解其关键模块,并探讨如何将其改造以适应更复杂的场景。

2. 构建核心:SAM预测器的深度交互封装

让我们先搭建一个比官方示例更清晰、更模块化的交互式脚本骨架。这个脚本的核心是封装SAM的预测过程,并提供一个可扩展的交互框架。

import cv2
import numpy as np
import os
from pathlib import Path
from typing import List, Tuple, Optional, Any
import torch

class SAMInteractiveSegmenter:
    """
    一个封装了SAM模型,并提供灵活交互逻辑的类。
    支持点提示、框提示的添加、撤销、预测与结果选择。
    """
    def __init__(self, model_type: str = "vit_h", checkpoint_path: str = "./sam_vit_h_4b8939.pth", device: str = "cuda"):
        """
        初始化SAM模型和预测器。
        
        参数:
            model_type: 模型架构,如 'vit_h', 'vit_l', 'vit_b'
            checkpoint_path: 模型权重文件路径
            device: 运行设备,'cuda' 或 'cpu'
        """
        from segment_anything import sam_model_registry, SamPredictor
        
        self.device = device if torch.cuda.is_available() and device == "cuda" else "cpu"
        print(f"正在加载SAM模型 ({model_type}) 到设备: {self.device}")
        
        sam = sam_model_registry[model_type](checkpoint=checkpoint_path)
        sam.to(device=self.device)
        self.predictor = SamPredictor(sam)
        
        # 状态管理
        self.input_points: List[List[int]] = []  # 存储点坐标 [x, y]
        self.input_labels: List[int] = []         # 存储点标签 1=前景, 0=背景
        self.input_box: Optional[List[int]] = None # 存储框坐标 [x1, y1, x2, y2]
        self.current_masks: Optional[np.ndarray] = None  # 当前预测的所有掩码 (N, H, W)
        self.current_scores: Optional[np.ndarray] = None # 对应的置信度分数
        self.selected_mask_idx: int = 0            # 当前选中的掩码索引
        self.current_logits: Optional[torch.Tensor] = None # 上一次预测的logits,用于迭代优化
        
    def set_image(self, image_bgr: np.ndarray):
        """设置当前要处理的图像(BGR格式)。"""
        image_rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)
        self.predictor.set_image(image_rgb)
        self._reset_interactive_state() # 每次设置新图像,重置交互状态
        print("图像已设置,可开始添加提示。")
    
    def _reset_interactive_state(self):
        """重置所有交互相关的状态。"""
        self.input_points.clear()
        self.input_labels.clear()
        self.input_box = None
        self.current_masks = None
        self.current_scores = None
        self.selected_mask_idx = 0
        self.current_logits = None
    
    def add_point(self, point: Tuple[int, int], is_positive: bool = True):
        """添加一个点提示。"""
        self.input_points.append(list(point))
        self.input_labels.append(1 if is_positive else 0)
        print(f"添加{'前景' if is_positive else '背景'}点: {point}")
    
    def add_box(self, box: Tuple[int, int, int, int]):
        """添加一个框提示。格式: (x_min, y_min, x_max, y_max)。"""
        self.input_box = list(box)
        print(f"添加框提示: {box}")
    
    def undo_last_point(self):
        """撤销最后一个添加的点提示。"""
        if self.input_points:
            removed_point = self.input_points.pop()
            removed_label = self.input_labels.pop()
            print(f"撤销点: {removed_point} (标签: {removed_label})")
        else:
            print("没有可撤销的点提示。")
    
    def predict(self, multimask_output: bool = True) -> bool:
        """
        基于当前所有提示(点、框)进行预测。
        
        返回:
            bool: 预测是否成功(至少有一个提示)。
        """
        if not (self.input_points or self.input_box):
            print("错误:预测需要至少一个点提示或框提示。")
            return False
        
        point_coords = np.array(self.input_points) if self.input_points else None
        point_labels = np.array(self.input_labels) if self.input_labels else None
        box = np.array(self.input_box).reshape(1, 4) if self.input_box is not None else None
        
        # 调用SAM核心预测函数
        masks, scores, logits = self.predictor.predict(
            point_coords=point_coords,
            point_labels=point_labels,
            box=box,
            mask_input=self.current_logits[None, :, :] if self.current_logits is not None else None,
            multimask_output=multimask_output,
        )
        
        self.current_masks = masks
        self.current_scores = scores
        self.current_logits = logits[np.argmax(scores), :, :] if multimask_output else logits[0, :, :]
        self.selected_mask_idx = np.argmax(scores) if multimask_output else 0
        
        print(f"预测完成。生成 {len(masks)} 个候选掩码。")
        if multimask_output:
            print(f"候选掩码置信度: {scores}")
            print(f"自动选择最高分掩码 (索引: {self.selected_mask_idx}, 分数: {scores[self.selected_mask_idx]:.3f})")
        return True

这个类SAMInteractiveSegmenter将SAM的交互逻辑状态化、模块化。它清晰地管理着提示信息、预测结果和中间状态。相比于直接操作全局变量和冗长的循环,这种面向对象的设计让后续的功能扩展(如保存、批量处理)变得清晰很多。

提示:在实际部署时,尤其是生产环境,建议将模型加载(__init__部分)与预测服务分离。可以将模型设为全局单例,避免每次处理图片都重复加载,这对提升吞吐量至关重要。

3. 掩码后处理:从二值矩阵到可用资产

得到二值掩码(0和1的矩阵)只是第一步。在计算机视觉流水线中,我们通常需要更丰富、更“可用”的输出。下面我们构建一个后处理工具箱,将原始的掩码转化为各种实用的形式。

3.1 基础工具函数

首先,定义一些核心的转换函数:

class MaskPostProcessor:
    """掩码后处理工具箱。"""
    
    @staticmethod
    def mask_to_alpha_channel(image_bgr: np.ndarray, mask: np.ndarray) -> np.ndarray:
        """
        将掩码应用于图像,生成带透明通道的PNG图像(RGBA格式)。
        前景区域不透明,背景区域完全透明。
        
        参数:
            image_bgr: 原始BGR图像
            mask: 二值掩码,形状 (H, W),值为0或1
        返回:
            rgba_image: RGBA格式的图像,形状 (H, W, 4)
        """
        # 确保mask是二值且类型正确
        binary_mask = (mask > 0).astype(np.uint8) * 255
        
        # 将BGR转为BGRA(先增加一个Alpha通道)
        bgra = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2BGRA)
        # 将掩码值赋给Alpha通道
        bgra[:, :, 3] = binary_mask
        # 也可以选择将背景区域的RGB值设为0,实现纯透明背景
        # bgra[mask == 0] = [0, 0, 0, 0]
        
        # 如果需要RGBA格式(某些库更常用),可以转换
        rgba = cv2.cvtColor(bgra, cv2.COLOR_BGRA2RGBA)
        return rgba
    
    @staticmethod
    def crop_to_mask_bbox(image: np.ndarray, mask: np.ndarray, padding: int = 5) -> Tuple[np.ndarray, np.ndarray]:
        """
        根据掩码的非零区域,裁剪图像和掩码,并可选添加边距。
        
        参数:
            image: 原始图像 (H, W, C)
            mask: 二值掩码 (H, W)
            padding: 在边界框四周添加的像素边距
        返回:
            cropped_image: 裁剪后的图像
            cropped_mask: 裁剪后的掩码
        """
        rows = np.any(mask, axis=1)
        cols = np.any(mask, axis=0)
        ymin, ymax = np.where(rows)[0][[0, -1]]
        xmin, xmax = np.where(cols)[0][[0, -1]]
        
        # 添加边距,并确保不超出图像边界
        h, w = image.shape[:2]
        ymin = max(0, ymin - padding)
        ymax = min(h, ymax + padding + 1)
        xmin = max(0, xmin - padding)
        xmax = min(w, xmax + padding + 1)
        
        cropped_image = image[ymin:ymax, xmin:xmax]
        cropped_mask = mask[ymin:ymax, xmin:xmax]
        return cropped_image, cropped_mask
    
    @staticmethod
    def extract_contours(mask: np.ndarray, simplify_epsilon: float = 0.005) -> List[np.ndarray]:
        """
        从掩码中提取轮廓(多边形点集)。
        使用OpenCV的findContours函数,并可选择对轮廓进行简化(Douglas-Peucker算法)。
        
        参数:
            mask: 二值掩码
            simplify_epsilon: 轮廓简化精度,值越大轮廓越简化
        返回:
            contours: 轮廓列表,每个轮廓是一个 (N, 1, 2) 的numpy数组
        """
        # 确保mask是8位单通道
        mask_uint8 = (mask * 255).astype(np.uint8)
        contours, _ = cv2.findContours(mask_uint8, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
        
        simplified_contours = []
        for cnt in contours:
            if simplify_epsilon > 0:
                perimeter = cv2.arcLength(cnt, True)
                epsilon = simplify_epsilon * perimeter
                approx = cv2.approxPolyDP(cnt, epsilon, True)
                simplified_contours.append(approx)
            else:
                simplified_contours.append(cnt)
        return simplified_contours

3.2 高级合成与可视化

除了基础转换,我们通常需要将分割结果以更直观的方式呈现,例如在原图上高亮显示,或者生成诊断图。

    @staticmethod
    def overlay_mask_on_image(image_bgr: np.ndarray, mask: np.ndarray, 
                              color: Tuple[int, int, int] = (0, 255, 0), 
                              alpha: float = 0.5) -> np.ndarray:
        """
        将掩码以半透明颜色叠加到原图上,用于可视化。
        
        参数:
            image_bgr: 背景图像
            mask: 二值掩码
            color: BGR格式的叠加颜色,默认绿色
            alpha: 叠加的透明度 (0-1)
        返回:
            overlayed_image: 叠加后的图像
        """
        overlay = image_bgr.copy()
        colored_mask = np.zeros_like(image_bgr)
        colored_mask[mask > 0] = color
        
        cv2.addWeighted(colored_mask, alpha, overlay, 1 - alpha, 0, overlay)
        # 可选:在边界处画上轮廓线,更清晰
        contours = MaskPostProcessor.extract_contours(mask, simplify_epsilon=0.002)
        cv2.drawContours(overlay, contours, -1, (255, 255, 255), thickness=2)
        return overlay
    
    @staticmethod
    def generate_diagnostic_figure(image_bgr: np.ndarray, masks: np.ndarray, scores: np.ndarray, 
                                   selected_idx: int) -> np.ndarray:
        """
        生成一个诊断图,在一个画布上并排显示原图、所有候选掩码叠加图以及选中的掩码。
        这对于调试和选择最佳掩码非常有用。
        
        参数:
            image_bgr: 原图
            masks: 所有候选掩码,形状 (N, H, W)
            scores: 对应的置信度分数
            selected_idx: 当前选中的掩码索引
        返回:
            diagnostic_img: 合成的诊断图像
        """
        h, w = image_bgr.shape[:2]
        # 创建一个大画布:一行三列 (原图,所有掩码热力图,选中掩码叠加图)
        canvas_h = h
        canvas_w = w * 3
        canvas = np.zeros((canvas_h, canvas_w, 3), dtype=np.uint8)
        
        # 第一列:原图
        canvas[0:h, 0:w] = image_bgr
        
        # 第二列:所有掩码叠加的热力图(用不同颜色/透明度)
        mask_overlay_all = image_bgr.copy()
        colors = [(0, 0, 255), (0, 255, 0), (255, 0, 0)]  # 红,绿,蓝
        for i, (mask, score) in enumerate(zip(masks, scores)):
            color = colors[i % len(colors)]
            alpha = 0.3 + 0.4 * (score)  # 分数越高越不透明
            colored_region = np.zeros_like(image_bgr)
            colored_region[mask > 0] = color
            cv2.addWeighted(colored_region, alpha, mask_overlay_all, 1, 0, mask_overlay_all)
            # 在图上标注分数
            cv2.putText(mask_overlay_all, f'M{i}:{score:.2f}', (20, 40 + i*30), 
                        cv2.FONT_HERSHEY_SIMPLEX, 0.7, color, 2)
        canvas[0:h, w:2*w] = mask_overlay_all
        
        # 第三列:选中的掩码清晰叠加
        selected_mask_overlay = MaskPostProcessor.overlay_mask_on_image(
            image_bgr, masks[selected_idx], color=(0, 200, 255), alpha=0.7
        )
        cv2.putText(selected_mask_overlay, f'Selected (Idx:{selected_idx})', (20, 40), 
                    cv2.FONT_HERSHEY_SIMPLEX, 0.8, (255, 255, 255), 2)
        canvas[0:h, 2*w:3*w] = selected_mask_overlay
        
        # 画分隔线
        cv2.line(canvas, (w, 0), (w, h), (255, 255, 255), 2)
        cv2.line(canvas, (2*w, 0), (2*w, h), (255, 255, 255), 2)
        
        return canvas

通过这个MaskPostProcessor类,我们将常见的后处理需求封装成了可复用的方法。在实际项目中,你可以像搭积木一样调用它们,组合出所需的输出格式。

4. 实战:构建一个批处理与结果管理系统

现在,我们将前两部分的模块组合起来,构建一个面向批量任务的脚本。这个脚本会遍历一个文件夹中的所有图片,为每张图片应用一套预定义或交互式的分割逻辑,并将多种格式的结果保存到结构化的目录中。

4.1 设计批处理流水线

我们设计一个BatchProcessingPipeline类来管理整个流程。它的核心思想是可配置和可扩展。

import json
from datetime import datetime

class BatchProcessingPipeline:
    """
    批量处理管道。
    支持自动模式(每张图使用固定提示)和交互式模式(人工逐张处理)。
    """
    
    def __init__(self, segmenter: SAMInteractiveSegmenter, 
                 input_dir: str, 
                 output_root: str):
        self.segmenter = segmenter
        self.input_dir = Path(input_dir)
        self.output_root = Path(output_root)
        self.output_root.mkdir(parents=True, exist_ok=True)
        
        # 定义输出子目录结构
        self.dirs = {
            'masks': self.output_root / 'masks_binary',      # 原始二值掩码
            'cropped': self.output_root / 'objects_cropped', # 裁剪后的对象图像
            'visualization': self.output_root / 'overlays',  # 可视化叠加图
            'metadata': self.output_root / 'metadata'        # 元数据(如轮廓、面积)
        }
        for d in self.dirs.values():
            d.mkdir(exist_ok=True)
        
        # 支持的图片格式
        self.image_extensions = {'.jpg', '.jpeg', '.png', '.bmp', '.tiff'}
    
    def get_image_list(self) -> List[Path]:
        """获取输入目录下所有支持的图像文件列表。"""
        image_files = []
        for ext in self.image_extensions:
            image_files.extend(self.input_dir.glob(f'*{ext}'))
            image_files.extend(self.input_dir.glob(f'*{ext.upper()}'))
        return sorted(image_files)  # 按文件名排序
    
    def process_single_image_auto(self, image_path: Path, 
                                  fixed_points: Optional[List[Tuple[int, int, bool]]] = None,
                                  fixed_box: Optional[Tuple[int, int, int, int]] = None) -> dict:
        """
        自动处理单张图像:使用预定义的提示点或框进行分割。
        适用于背景简单、目标位置固定的场景(如产品图)。
        
        参数:
            image_path: 图像路径
            fixed_points: 预定义的点列表,每个元素为 (x, y, is_positive)
            fixed_box: 预定义的框 (x1, y1, x2, y2)
        返回:
            result_info: 包含处理结果和元数据的字典
        """
        result_info = {'image_name': image_path.name, 'success': False}
        image = cv2.imread(str(image_path))
        if image is None:
            print(f"无法读取图像: {image_path}")
            return result_info
        
        self.segmenter.set_image(image)
        
        # 应用预定义提示
        if fixed_points:
            for x, y, is_pos in fixed_points:
                self.segmenter.add_point((x, y), is_positive=is_pos)
        if fixed_box:
            self.segmenter.add_box(fixed_box)
        
        # 执行预测
        if not self.segmenter.predict(multimask_output=True):
            print(f"对 {image_path.name} 预测失败,无有效提示。")
            return result_info
        
        # 获取最佳掩码
        best_mask = self.segmenter.current_masks[self.segmenter.selected_mask_idx]
        best_score = self.segmenter.current_scores[self.segmenter.selected_mask_idx]
        
        # 生成并保存各种输出
        base_name = image_path.stem
        timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
        
        # 1. 保存原始二值掩码
        mask_filename = self.dirs['masks'] / f"{base_name}_{timestamp}_mask.png"
        cv2.imwrite(str(mask_filename), (best_mask * 255).astype(np.uint8))
        
        # 2. 保存裁剪后的对象(带透明背景)
        cropped_img, cropped_mask = MaskPostProcessor.crop_to_mask_bbox(image, best_mask, padding=10)
        rgba_cropped = MaskPostProcessor.mask_to_alpha_channel(cropped_img, cropped_mask)
        cropped_filename = self.dirs['cropped'] / f"{base_name}_{timestamp}_obj.png"
        cv2.imwrite(str(cropped_filename), cv2.cvtColor(rgba_cropped, cv2.COLOR_RGBA2BGRA))
        
        # 3. 保存可视化叠加图
        overlay = MaskPostProcessor.overlay_mask_on_image(image, best_mask, color=(0, 200, 255), alpha=0.6)
        overlay_filename = self.dirs['visualization'] / f"{base_name}_{timestamp}_overlay.jpg"
        cv2.imwrite(str(overlay_filename), overlay)
        
        # 4. 计算并保存元数据
        contours = MaskPostProcessor.extract_contours(best_mask)
        area_pixels = np.sum(best_mask)
        metadata = {
            'image': image_path.name,
            'mask_file': mask_filename.name,
            'score': float(best_score),
            'area_pixels': int(area_pixels),
            'contour_count': len(contours),
            'processing_time': timestamp,
            'prompts_used': {
                'points': fixed_points if fixed_points else [],
                'box': fixed_box if fixed_box else None
            }
        }
        meta_filename = self.dirs['metadata'] / f"{base_name}_{timestamp}_meta.json"
        with open(meta_filename, 'w') as f:
            json.dump(metadata, f, indent=2)
        
        result_info.update({
            'success': True,
            'mask_path': mask_filename,
            'cropped_path': cropped_filename,
            'overlay_path': overlay_filename,
            'metadata_path': meta_filename,
            'score': best_score,
            'area': area_pixels
        })
        print(f"处理完成: {image_path.name} -> 分数: {best_score:.3f}, 面积: {area_pixels} 像素")
        return result_info

4.2 交互式批处理与状态保存

对于需要人工干预的复杂图片,我们可以设计一个交互式批处理循环,并加入状态保存功能,防止程序意外中断导致工作丢失。

    def run_interactive_session(self, start_index: int = 0):
        """
        运行交互式批处理会话。
        逐张显示图片,允许用户通过鼠标和键盘添加提示、预测、选择并保存结果。
        处理状态会自动保存,以便中断后恢复。
        """
        image_files = self.get_image_list()
        if not image_files:
            print("输入目录中没有找到图像文件。")
            return
        
        # 尝试加载之前的处理进度
        state_file = self.output_root / 'processing_state.json'
        processed = self._load_processing_state(state_file)
        
        cv2.namedWindow("SAM Interactive Batch")
        print("\n=== 交互式批处理控制说明 ===")
        print("鼠标左键: 添加前景点")
        print("鼠标右键: 添加背景点")
        print("按 'b': 进入框标注模式 (拖动鼠标绘制框)")
        print("按 'w': 执行预测")
        print("按 'a'/'d': 在多个预测掩码间切换")
        print("按 's': 保存当前选中掩码的所有输出")
        print("按 'n': 处理下一张图")
        print("按 'p': 返回上一张图")
        print("按 'z': 撤销最后一个点")
        print("按 'c': 清除所有当前提示")
        print("按 'q': 退出程序")
        print("============================\n")
        
        current_idx = start_index
        drawing_box = False
        box_start = None
        
        def mouse_callback(event, x, y, flags, param):
            nonlocal drawing_box, box_start
            if event == cv2.EVENT_LBUTTONDOWN:
                if drawing_box:
                    box_start = (x, y)
                else:
                    self.segmenter.add_point((x, y), is_positive=True)
            elif event == cv2.EVENT_RBUTTONDOWN and not drawing_box:
                self.segmenter.add_point((x, y), is_positive=False)
            elif event == cv2.EVENT_MOUSEMOVE and drawing_box and box_start:
                # 实时绘制选框
                img_disp = param['image'].copy()
                cv2.rectangle(img_disp, box_start, (x, y), (255, 0, 0), 2)
                cv2.imshow("SAM Interactive Batch", img_disp)
            elif event == cv2.EVENT_LBUTTONUP and drawing_box and box_start:
                box_end = (x, y)
                x1, y1 = min(box_start[0], box_end[0]), min(box_start[1], box_end[1])
                x2, y2 = max(box_start[0], box_end[0]), max(box_start[1], box_end[1])
                if abs(x2 - x1) > 5 and abs(y2 - y1) > 5:  # 避免误触
                    self.segmenter.add_box((x1, y1, x2, y2))
                drawing_box = False
                box_start = None
        
        for idx in range(current_idx, len(image_files)):
            img_path = image_files[idx]
            if str(img_path) in processed:
                print(f"跳过已处理的图片: {img_path.name}")
                continue
                
            print(f"\n--- 处理第 {idx+1}/{len(image_files)} 张: {img_path.name} ---")
            image = cv2.imread(str(img_path))
            if image is None:
                continue
                
            self.segmenter.set_image(image)
            cv2.setMouseCallback("SAM Interactive Batch", mouse_callback, {'image': image})
            
            while True:
                # 绘制当前状态
                display_img = image.copy()
                # 绘制已有点
                for pt, label in zip(self.segmenter.input_points, self.segmenter.input_labels):
                    color = (0, 255, 0) if label == 1 else (0, 0, 255)
                    cv2.circle(display_img, tuple(pt), 5, color, -1)
                # 绘制已选框
                if self.segmenter.input_box:
                    x1, y1, x2, y2 = self.segmenter.input_box
                    cv2.rectangle(display_img, (x1, y1), (x2, y2), (255, 0, 0), 2)
                # 如果已有预测结果,绘制当前选中的掩码
                if self.segmenter.current_masks is not None:
                    overlay = MaskPostProcessor.overlay_mask_on_image(
                        display_img, 
                        self.segmenter.current_masks[self.segmenter.selected_mask_idx],
                        color=(0, 200, 255), 
                        alpha=0.5
                    )
                    display_img = overlay
                    # 显示当前掩码信息
                    info = f"Mask {self.segmenter.selected_mask_idx+1}/{len(self.segmenter.current_masks)} | Score: {self.segmenter.current_scores[self.segmenter.selected_mask_idx]:.3f}"
                    cv2.putText(display_img, info, (20, 40), cv2.FONT_HERSHEY_SIMPLEX, 0.7, (255, 255, 255), 2)
                
                # 显示控制说明
                status_text = [
                    f"Image: {img_path.name} ({idx+1}/{len(image_files)})",
                    "L-Click: Foreground | R-Click: Background | 'b': Draw Box | 'w': Predict",
                    "'a'/'d': Switch Mask | 's': Save | 'n': Next | 'p': Prev | 'z': Undo | 'c': Clear | 'q': Quit"
                ]
                y_offset = 70
                for line in status_text:
                    cv2.putText(display_img, line, (20, y_offset), cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0, 255, 255), 1)
                    y_offset += 25
                
                cv2.imshow("SAM Interactive Batch", display_img)
                key = cv2.waitKey(1) & 0xFF
                
                if key == ord('b'):
                    drawing_box = True
                    print("进入框标注模式,请拖动鼠标绘制选框。")
                elif key == ord('w'):
                    if self.segmenter.predict():
                        print("预测完成。使用 'a'/'d' 键切换查看不同候选掩码。")
                elif key == ord('a') and self.segmenter.current_masks is not None:
                    self.segmenter.selected_mask_idx = (self.segmenter.selected_mask_idx - 1) % len(self.segmenter.current_masks)
                elif key == ord('d') and self.segmenter.current_masks is not None:
                    self.segmenter.selected_mask_idx = (self.segmenter.selected_mask_idx + 1) % len(self.segmenter.current_masks)
                elif key == ord('s') and self.segmenter.current_masks is not None:
                    # 保存当前结果
                    result = self.process_single_image_auto(img_path)  # 复用自动处理函数保存
                    if result['success']:
                        processed.add(str(img_path))
                        self._save_processing_state(state_file, processed)
                        print(f"结果已保存。")
                    break
                elif key == ord('n'):
                    # 保存当前图片的处理状态(即使没保存结果)
                    processed.add(str(img_path))
                    self._save_processing_state(state_file, processed)
                    break
                elif key == ord('p'):
                    # 回到上一张,但至少是第0张
                    idx = max(idx - 2, -1)  # 因为循环会+1,所以减2
                    break
                elif key == ord('z'):
                    self.segmenter.undo_last_point()
                elif key == ord('c'):
                    self.segmenter._reset_interactive_state()
                    print("所有提示已清除。")
                elif key == ord('q'):
                    print("用户中断处理。")
                    cv2.destroyAllWindows()
                    return
            
            if key == ord('q'):
                break
        
        cv2.destroyAllWindows()
        print(f"\n批处理会话结束。已处理 {len(processed)} 张图片。")
    
    def _load_processing_state(self, state_file: Path) -> set:
        """从JSON文件加载已处理图片的集合。"""
        if state_file.exists():
            try:
                with open(state_file, 'r') as f:
                    return set(json.load(f))
            except:
                pass
        return set()
    
    def _save_processing_state(self, state_file: Path, processed_set: set):
        """将已处理图片的集合保存到JSON文件。"""
        with open(state_file, 'w') as f:
            json.dump(list(processed_set), f)

这个批处理管道展示了如何将交互逻辑、状态管理、文件IO和结果生成紧密结合。它提供了两种模式:全自动的process_single_image_auto适用于规则任务;run_interactive_session则提供了强大的手动控制能力,并贴心地加入了进度保存功能。

5. 性能优化与高级技巧

当处理高分辨率图像或大规模数据集时,性能成为关键。以下是一些提升脚本效率的实战技巧。

5.1 图像预处理与提示优化

SAM模型对输入图像尺寸敏感。直接使用超大原图会显著增加内存消耗和推理时间。

def preprocess_image_for_sam(image_bgr: np.ndarray, max_long_side: int = 1024) -> np.ndarray:
    """
    将图像缩放到适合SAM处理的尺寸,保持长宽比。
    SAM在长边1024左右的图像上表现和效率较好。
    
    参数:
        image_bgr: 输入BGR图像
        max_long_side: 目标长边的最大像素值
    返回:
        resized_image: 缩放后的图像
        scale_factor: 缩放比例,用于将提示坐标映射回原图
    """
    h, w = image_bgr.shape[:2]
    scale = max_long_side / max(h, w)
    if scale < 1.0:
        new_w = int(w * scale)
        new_h = int(h * scale)
        resized = cv2.resize(image_bgr, (new_w, new_h), interpolation=cv2.INTER_LINEAR)
        return resized, scale
    else:
        return image_bgr.copy(), 1.0

# 在SAMInteractiveSegmenter的set_image方法中集成预处理
def set_image_with_preprocess(self, image_bgr: np.ndarray, max_long_side: int = 1024):
    """设置图像,并自动进行预处理缩放。"""
    self.original_image = image_bgr
    self.scale_factor = 1.0
    
    processed_img, scale = preprocess_image_for_sam(image_bgr, max_long_side)
    self.scale_factor = scale
    # 注意:需要将后续添加的提示点坐标也按此比例缩放
    # 可以在add_point等方法内部处理,或者存储原始坐标,在predict前统一转换。
    self.set_image(processed_img)  # 调用原始的set_image处理缩放后的图

5.2 利用multimask_output进行迭代优化

SAM的multimask_output=True会返回三个候选掩码和对应的logits。我们可以利用分数最高的logits作为下一次预测的mask_input,进行迭代细化,这在物体边缘复杂时特别有效。

def iterative_refinement(self, image_bgr: np.ndarray, initial_points: List[Tuple], 
                         refinement_steps: int = 2):
    """
    迭代优化分割结果。
    1. 用初始点进行预测,得到最佳掩码和logits。
    2. 将最佳logits作为mask_input,结合可能的新提示点,再次预测。
    3. 重复,逐步细化。
    """
    self.set_image(image_bgr)
    for point in initial_points:
        self.add_point((point[0], point[1]), is_positive=point[2])
    
    all_masks = []
    all_scores = []
    
    for step in range(refinement_steps):
        success = self.predict(multimask_output=True)
        if not success:
            break
        best_idx = self.selected_mask_idx
        all_masks.append(self.current_masks[best_idx])
        all_scores.append(self.current_scores[best_idx])
        
        print(f"Refinement step {step+1}, best score: {self.current_scores[best_idx]:.3f}")
        
        if step < refinement_steps - 1:
            # 在实际应用中,这里可以加入一些启发式规则,
            # 例如在掩码边界不确定性高的区域自动添加负点,或者由用户交互添加。
            # 本次迭代的logits已自动保存在self.current_logits中,供下一次predict使用。
            # 模拟:在掩码边缘随机采样一些点作为负点(示例,实际需更智能)
            pass 
    
    # 返回所有迭代步骤中最好的那个掩码
    best_overall_idx = np.argmax(all_scores) if all_scores else 0
    return all_masks[best_overall_idx], all_scores[best_overall_idx]

5.3 结果缓存与异步处理

对于批量任务,尤其是自动模式,可以缓存已编码的图像特征,避免对同一张图片重复运行predictor.set_image()。

from functools import lru_cache
import hashlib

class CachedSAMSegmenter(SAMInteractiveSegmenter):
    """带图像特征缓存的SAM分割器。"""
    
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        self._image_cache = {}
        
    def _get_image_hash(self, image_array: np.ndarray) -> str:
        """生成图像的哈希值作为缓存键。"""
        # 使用分辨率和小型缩略图的哈希,平衡速度与准确性
        small_img = cv2.resize(image_array, (64, 64))
        return hashlib.md5(small_img.tobytes()).hexdigest()
    
    def set_image(self, image_bgr: np.ndarray):
        """重写set_image,加入缓存逻辑。"""
        image_rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB)
        img_hash = self._get_image_hash(image_rgb)
        
        if img_hash not in self._image_cache:
            # 未缓存,调用原始方法并缓存特征
            self.predictor.set_image(image_rgb)
            # 注意:这里我们缓存的是predictor的内部状态,实际中SAM的predictor
            # 在set_image后会计算并存储图像编码。我们这里简化演示缓存思想。
            # 更复杂的实现可能需要直接缓存`predictor.features`等属性。
            self._image_cache[img_hash] = {
                'original_shape': image_bgr.shape,
                # 在实际中,可能需要根据SAM版本保存特定的特征张量
            }
            print(f"图像特征未命中缓存,已计算并缓存。")
        else:
            print(f"图像特征命中缓存,跳过编码计算。")
            # 这里需要将predictor的内部状态恢复为缓存的值。
            # 由于SAM的predictor不直接提供该接口,此示例仅说明思路。
            # 一种替代方案是,对于完全相同的图片,跳过set_image,但这要求图片完全未变。
            # 对于批量处理同一张图不同区域的情况,缓存能极大提升速度。
        
        # 无论缓存与否,都需要重置交互状态
        self._reset_interactive_state()

将这些优化技巧融入之前的批处理管道,可以显著提升处理效率,尤其是在处理大量相似图片或需要反复对同一图片尝试不同提示时。

6. 错误处理与日志记录

一个健壮的脚本必须能妥善处理异常,并记录详细日志以供调试。我们在核心类中加入这些能力。

import logging
import sys
import traceback

def setup_logging(output_dir: Path):
    """配置日志,同时输出到控制台和文件。"""
    log_file = output_dir / 'sam_processing.log'
    logging.basicConfig(
        level=logging.INFO,
        format='%(asctime)s - %(name)s - %(levelname)s - %(message)s',
        handlers=[
            logging.FileHandler(log_file, encoding='utf-8'),
            logging.StreamHandler(sys.stdout)
        ]
    )
    return logging.getLogger(__name__)

# 在BatchProcessingPipeline的__init__中初始化日志
self.logger = setup_logging(self.output_root)

# 在process_single_image_auto等方法中,用try-except包裹核心逻辑
try:
    # ... 处理逻辑 ...
    self.logger.info(f"成功处理图像: {image_path.name}, 分数: {best_score:.3f}")
except FileNotFoundError as e:
    self.logger.error(f"文件未找到错误: {image_path} - {e}")
    result_info['error'] = str(e)
except RuntimeError as e:
    self.logger.error(f"运行时错误 (可能是CUDA内存不足): {image_path} - {e}")
    # 可以尝试降级模型或切换到CPU模式
    result_info['error'] = str(e)
except Exception as e:
    self.logger.error(f"处理图像 {image_path} 时发生未知错误: {e}")
    self.logger.error(traceback.format_exc())
    result_info['error'] = str(e)

通过系统的错误处理和日志记录,当在服务器上运行 overnight 的批处理任务时,你可以轻松定位是哪张图片、哪个步骤出了问题,是模型加载失败、显存溢出,还是文件权限问题。

7. 整合与部署:从脚本到工具

最后,我们将所有模块整合到一个主函数中,并提供一个简单的命令行接口,使其成为一个可以直接使用的工具。

# main.py
import argparse

def main():
    parser = argparse.ArgumentParser(description='SAM自定义图像分割与批处理工具')
    parser.add_argument('--input', type=str, required=True, help='输入图片目录路径')
    parser.add_argument('--output', type=str, required=True, help='输出根目录路径')
    parser.add_argument('--model-type', type=str, default='vit_h', choices=['vit_h', 'vit_l', 'vit_b'], help='SAM模型类型')
    parser.add_argument('--checkpoint', type=str, required=True, help='SAM模型权重文件路径')
    parser.add_argument('--mode', type=str, default='interactive', choices=['auto', 'interactive', 'batch'], 
                        help='运行模式: auto(自动), interactive(交互), batch(无交互批处理)')
    parser.add_argument('--device', type=str, default='cuda', help='计算设备,cuda 或 cpu')
    parser.add_argument('--fixed-points', type=str, default='', 
                        help='自动模式下的固定点提示,格式: x1,y1,label1;x2,y2,label2 (label:1前景,0背景)')
    
    args = parser.parse_args()
    
    # 1. 初始化分割器
    segmenter = SAMInteractiveSegmenter(
        model_type=args.model_type,
        checkpoint_path=args.checkpoint,
        device=args.device
    )
    
    # 2. 初始化处理管道
    pipeline = BatchProcessingPipeline(segmenter, args.input, args.output)
    
    # 3. 根据模式运行
    if args.mode == 'auto':
        # 解析固定点
        points = []
        if args.fixed_points:
            for pt_str in args.fixed_points.split(';'):
                x, y, label = map(int, pt_str.split(','))
                points.append((x, y, bool(label)))
        print(f"自动模式启动,使用固定点提示: {points}")
        image_files = pipeline.get_image_list()
        for img_path in image_files:
            pipeline.process_single_image_auto(img_path, fixed_points=points)
        print("自动批处理完成。")
        
    elif args.mode == 'interactive':
        print("启动交互式模式...")
        pipeline.run_interactive_session()
        
    elif args.mode == 'batch':
        # 无交互批处理,遍历所有图片,每张图都需要人工预先定义好提示(可通过外部配置文件)
        # 这里作为示例,假设每张图都用图像中心点作为前景点提示
        print("启动无交互批处理模式(使用图像中心作为默认点)...")
        image_files = pipeline.get_image_list()
        for img_path in image_files:
            image = cv2.imread(str(img_path))
            if image is None:
                continue
            h, w = image.shape[:2]
            center_point = [(w//2, h//2, True)]  # 中心点作为前景提示
            pipeline.process_single_image_auto(img_path, fixed_points=center_point)
        print("批处理完成。")
    
if __name__ == "__main__":
    main()

现在,你可以通过命令行灵活地调用这个工具了:

# 交互式处理一个文件夹的图片
python main.py --input ./my_images --output ./results --checkpoint ./sam_vit_h_4b8939.pth --mode interactive

# 自动模式,对所有图片使用同一个预定义点进行分割
python main.py --input ./product_shots --output ./auto_results --checkpoint ./sam_vit_b_01ec64.pth --mode auto --fixed-points "300,400,1;350,420,0"

# 无交互批处理(示例:每张图用中心点)
python main.py --input ./dataset --output ./batch_results --model-type vit_l --checkpoint ./sam_vit_l_0b3195.pth --mode batch

通过这样层层递进的构建,我们从最初的一个简单交互脚本,发展出了一个功能全面、结构清晰、易于扩展和部署的图像分割工具链。它解决了官方脚本在灵活性、输出格式和流程控制上的不足,让你能够真正将SAM的强大能力融入到自己特定的项目和需求中去。记住,这里的每一段代码都不是终点,而是一个起点,你可以根据实际遇到的具体问题,继续优化交互逻辑、增加后处理算法(如掩码平滑、小物体过滤),或者将其封装为REST API服务。

Logo

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

更多推荐