import os import cv2 import tempfile import spaces import gradio as gr import numpy as np import torch import matplotlib import matplotlib.pyplot as plt from PIL import Image, ImageDraw from typing import Iterable from transformers import ( Sam3Model, Sam3Processor, Sam3VideoModel, Sam3VideoProcessor, Sam3TrackerModel, Sam3TrackerProcessor ) # 使用 SAM 3.1 最新模型 MODEL_ID = "facebook/sam3.1" device = "cuda" if torch.cuda.is_available() else "cpu" print(f"Using compute device: {device}") print("Loading SAM 3.1 Models...") try: # 1. 图片文本分割模型 print(" ... Loading Image Text Model") IMG_MODEL = Sam3Model.from_pretrained(MODEL_ID).to(device) IMG_PROCESSOR = Sam3Processor.from_pretrained(MODEL_ID) # 2. 图片点击追踪模型 print(" ... Loading Image Tracker Model") TRK_MODEL = Sam3TrackerModel.from_pretrained(MODEL_ID).to(device) TRK_PROCESSOR = Sam3TrackerProcessor.from_pretrained(MODEL_ID) # 3. 视频分割模型 print(" ... Loading Video Model") VID_MODEL = Sam3VideoModel.from_pretrained(MODEL_ID).to(device, dtype=torch.bfloat16) VID_PROCESSOR = Sam3VideoProcessor.from_pretrained(MODEL_ID) print("All Models loaded successfully!") except Exception as e: print(f"CRITICAL ERROR LOADING MODELS: {e}") IMG_MODEL = None IMG_PROCESSOR = None TRK_MODEL = None TRK_PROCESSOR = None VID_MODEL = None VID_PROCESSOR = None # --- 工具函数 --- def apply_mask_overlay(base_image, mask_data, opacity=0.5): """在图片上绘制分割遮罩""" if isinstance(base_image, np.ndarray): base_image = Image.fromarray(base_image) base_image = base_image.convert("RGBA") if mask_data is None or len(mask_data) == 0: return base_image.convert("RGB") if isinstance(mask_data, torch.Tensor): mask_data = mask_data.cpu().numpy() mask_data = mask_data.astype(np.uint8) if mask_data.ndim == 4: mask_data = mask_data[0] if mask_data.ndim == 3 and mask_data.shape[0] == 1: mask_data = mask_data[0] num_masks = mask_data.shape[0] if mask_data.ndim == 3 else 1 if mask_data.ndim == 2: mask_data = [mask_data] num_masks = 1 try: color_map = matplotlib.colormaps["rainbow"].resampled(max(num_masks, 1)) except AttributeError: import matplotlib.cm as cm color_map = cm.get_cmap("rainbow").resampled(max(num_masks, 1)) rgb_colors = [tuple(int(c * 255) for c in color_map(i)[:3]) for i in range(num_masks)] composite_layer = Image.new("RGBA", base_image.size, (0, 0, 0, 0)) for i, single_mask in enumerate(mask_data): mask_bitmap = Image.fromarray((single_mask * 255).astype(np.uint8)) if mask_bitmap.size != base_image.size: mask_bitmap = mask_bitmap.resize(base_image.size, resample=Image.NEAREST) fill_color = rgb_colors[i] color_fill = Image.new("RGBA", base_image.size, fill_color + (0,)) mask_alpha = mask_bitmap.point(lambda v: int(v * opacity) if v > 0 else 0) color_fill.putalpha(mask_alpha) composite_layer = Image.alpha_composite(composite_layer, color_fill) return Image.alpha_composite(base_image, composite_layer).convert("RGB") def draw_points_on_image(image, points): """在图片上绘制红色标记点""" if isinstance(image, np.ndarray): image = Image.fromarray(image) draw_img = image.copy() draw = ImageDraw.Draw(draw_img) for pt in points: x, y = pt r = 8 draw.ellipse((x-r, y-r, x+r, y+r), fill="red", outline="white", width=4) return draw_img @spaces.GPU def run_image_segmentation(source_img, text_query, conf_thresh=0.5): if IMG_MODEL is None or IMG_PROCESSOR is None: raise gr.Error("模型加载失败,请检查日志。") if source_img is None or not text_query: raise gr.Error("请提供图片和文本提示。") try: pil_image = source_img.convert("RGB") model_inputs = IMG_PROCESSOR(images=pil_image, text=text_query, return_tensors="pt").to(device) with torch.no_grad(): inference_output = IMG_MODEL(**model_inputs) processed_results = IMG_PROCESSOR.post_process_instance_segmentation( inference_output, threshold=conf_thresh, mask_threshold=0.5, target_sizes=model_inputs.get("original_sizes").tolist() )[0] annotation_list = [] raw_masks = processed_results['masks'].cpu().numpy() raw_scores = processed_results['scores'].cpu().numpy() for idx, mask_array in enumerate(raw_masks): label_str = f"{text_query} ({raw_scores[idx]:.2f})" annotation_list.append((mask_array, label_str)) return (pil_image, annotation_list) except Exception as e: raise gr.Error(f"图片处理出错: {e}") @spaces.GPU def run_image_click_gpu(input_image, x, y, points_state, labels_state): if TRK_MODEL is None or TRK_PROCESSOR is None: raise gr.Error("追踪模型加载失败。") if input_image is None: return input_image, [], [] if points_state is None: points_state = []; labels_state = [] points_state.append([x, y]) labels_state.append(1) try: input_points = [[points_state]] input_labels = [[labels_state]] inputs = TRK_PROCESSOR(images=input_image, input_points=input_points, input_labels=input_labels, return_tensors="pt").to(device) with torch.no_grad(): outputs = TRK_MODEL(**inputs, multimask_output=False) masks = TRK_PROCESSOR.post_process_masks(outputs.pred_masks.cpu(), inputs["original_sizes"], binarize=True)[0] final_img = apply_mask_overlay(input_image, masks[0]) final_img = draw_points_on_image(final_img, points_state) return final_img, points_state, labels_state except Exception as e: print(f"Tracker Error: {e}") return input_image, points_state, labels_state def image_click_handler(image, evt: gr.SelectData, points_state, labels_state): x, y = evt.index return run_image_click_gpu(image, x, y, points_state, labels_state) def calc_timeout_duration(vid_file, *args): return args[-1] if args else 60 @spaces.GPU(duration=calc_timeout_duration) def run_video_segmentation(source_vid, text_query, frame_limit, time_limit): if VID_MODEL is None or VID_PROCESSOR is None: raise gr.Error("视频模型加载失败。") if not source_vid or not text_query: raise gr.Error("请提供视频和文本提示。") try: video_cap = cv2.VideoCapture(source_vid) vid_fps = video_cap.get(cv2.CAP_PROP_FPS) vid_w = int(video_cap.get(cv2.CAP_PROP_FRAME_WIDTH)) vid_h = int(video_cap.get(cv2.CAP_PROP_FRAME_HEIGHT)) video_frames = [] counter = 0 while video_cap.isOpened(): ret, frame = video_cap.read() if not ret or (frame_limit > 0 and counter >= frame_limit): break video_frames.append(cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)) counter += 1 video_cap.release() session = VID_PROCESSOR.init_video_session(video=video_frames, inference_device=device, dtype=torch.bfloat16) session = VID_PROCESSOR.add_text_prompt(inference_session=session, text=text_query) temp_out_path = tempfile.mktemp(suffix=".mp4") video_writer = cv2.VideoWriter(temp_out_path, cv2.VideoWriter_fourcc(*'mp4v'), vid_fps, (vid_w, vid_h)) for model_out in VID_MODEL.propagate_in_video_iterator(inference_session=session, max_frame_num_to_track=len(video_frames)): post_processed = VID_PROCESSOR.postprocess_outputs(session, model_out) f_idx = model_out.frame_idx original_pil = Image.fromarray(video_frames[f_idx]) if 'masks' in post_processed: detected_masks = post_processed['masks'] if detected_masks.ndim == 4: detected_masks = detected_masks.squeeze(1) final_frame = apply_mask_overlay(original_pil, detected_masks) else: final_frame = original_pil video_writer.write(cv2.cvtColor(np.array(final_frame), cv2.COLOR_RGB2BGR)) video_writer.release() return temp_out_path, "视频处理完成 ✅" except Exception as e: return None, f"视频处理出错: {str(e)}" custom_css = """ #col-container { margin: 0 auto; max-width: 1200px; } #main-title h1 { font-size: 2.2em !important; } """ with gr.Blocks() as demo: with gr.Column(elem_id="col-container"): gr.Markdown("# **SAM 3.1: Segment Anything with Concepts**", elem_id="main-title") gr.Markdown("使用 **SAM 3.1** 通过文本提示或交互式点击,分割图片和视频中的物体。") with gr.Tabs(): with gr.Tab("图片文本分割"): with gr.Row(): with gr.Column(scale=1): image_input = gr.Image(label="上传图片", type="pil", height=350) txt_prompt_img = gr.Textbox(label="文本提示", placeholder="例如: cat, face, car wheel") with gr.Accordion("高级设置", open=False): conf_slider = gr.Slider(0.0, 1.0, value=0.45, step=0.05, label="置信度阈值") btn_process_img = gr.Button("开始分割", variant="primary") with gr.Column(scale=1.5): image_result = gr.AnnotatedImage(label="分割结果", height=410) btn_process_img.click( fn=run_image_segmentation, inputs=[image_input, txt_prompt_img, conf_slider], outputs=[image_result] ) with gr.Tab("视频分割"): with gr.Row(): with gr.Column(): video_input = gr.Video(label="上传视频", format="mp4", height=320) txt_prompt_vid = gr.Textbox(label="文本提示", placeholder="例如: person running, red car") with gr.Row(): frame_limiter = gr.Slider(10, 500, value=60, step=10, label="最大帧数") time_limiter = gr.Radio([60, 120, 180], value=60, label="超时时间(秒)") btn_process_vid = gr.Button("开始分割", variant="primary") with gr.Column(): video_result = gr.Video(label="处理结果") process_status = gr.Textbox(label="状态", interactive=False) btn_process_vid.click( run_video_segmentation, inputs=[video_input, txt_prompt_vid, frame_limiter, time_limiter], outputs=[video_result, process_status] ) with gr.Tab("图片点击分割"): with gr.Row(): with gr.Column(scale=1): img_click_input = gr.Image(type="pil", label="上传图片(点击要分割的物体)", interactive=True, height=450) with gr.Row(): img_click_clear = gr.Button("清除标记点", variant="primary") st_click_points = gr.State([]) st_click_labels = gr.State([]) with gr.Column(scale=1): img_click_output = gr.Image(type="pil", label="分割预览", height=450, interactive=False) img_click_input.select( image_click_handler, inputs=[img_click_input, st_click_points, st_click_labels], outputs=[img_click_output, st_click_points, st_click_labels] ) img_click_clear.click( lambda: (None, [], []), outputs=[img_click_output, st_click_points, st_click_labels] ) if __name__ == "__main__": demo.launch(css=custom_css, theme=gr.themes.Soft(), ssr_mode=False, mcp_server=True, show_error=True)