SAM模型实战:如何用自定义脚本实现更灵活的图像分割(Python代码详解)
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服务。
DAMO开发者矩阵,由阿里巴巴达摩院和中国互联网协会联合发起,致力于探讨最前沿的技术趋势与应用成果,搭建高质量的交流与分享平台,推动技术创新与产业应用链接,围绕“人工智能与新型计算”构建开放共享的开发者生态。
更多推荐

所有评论(0)