""" SegFormer-B0 最小化Demo 用于快速上手图像语义分割 """ import os import torch import numpy as np from PIL import Image import matplotlib.pyplot as plt from transformers import SegformerImageProcessor, SegformerForSemanticSegmentation class SegFormerDemo: def __init__(self, model_name="nvidia/segformer-b0-finetuned-ade-512-512"): """ 初始化SegFormer-B0模型 Args: model_name: 预训练模型名称 - "nvidia/segformer-b0-finetuned-ade-512-512" (ADE20K数据集,150类) - "nvidia/segformer-b0-finetuned-cityscapes-1024-1024" (城市街景) """ print(f"加载SegFormer模型: {model_name}") self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"使用设备: {self.device}") # 加载处理器和模型 self.processor = SegformerImageProcessor.from_pretrained(model_name) self.model = SegformerForSemanticSegmentation.from_pretrained(model_name) self.model.to(self.device) self.model.eval() # ADE20K数据集的类别数 self.num_classes = self.model.config.num_labels print(f"模型加载完成! 支持 {self.num_classes} 个类别") def segment_image(self, image_path): """ 对图像进行语义分割 Args: image_path: 图像路径 Returns: pred_seg: 分割结果 (H, W) 每个像素的类别ID original_image: 原始图像 """ # 读取图像 image = Image.open(image_path).convert("RGB") # 预处理 inputs = self.processor(images=image, return_tensors="pt") inputs = {k: v.to(self.device) for k, v in inputs.items()} # 推理 with torch.no_grad(): outputs = self.model(**inputs) logits = outputs.logits # 上采样到原始尺寸 upsampled_logits = torch.nn.functional.interpolate( logits, size=image.size[::-1], # (height, width) mode="bilinear", align_corners=False ) # 获取每个像素的类别 pred_seg = upsampled_logits.argmax(dim=1)[0].cpu().numpy() return pred_seg, np.array(image) def visualize_segmentation(self, image, pred_seg, save_path=None): """ 可视化分割结果 Args: image: 原始图像 (H, W, 3) pred_seg: 分割结果 (H, W) save_path: 保存路径 """ fig, axes = plt.subplots(1, 3, figsize=(15, 5)) # 原始图像 axes[0].imshow(image) axes[0].set_title("Original Image") axes[0].axis("off") # 分割mask(彩色) axes[1].imshow(pred_seg, cmap='tab20b') axes[1].set_title("Segmentation Mask") axes[1].axis("off") # 叠加显示 axes[2].imshow(image) axes[2].imshow(pred_seg, alpha=0.5, cmap='tab20b') axes[2].set_title("Overlay") axes[2].axis("off") plt.tight_layout() if save_path: plt.savefig(save_path, dpi=150, bbox_inches='tight') print(f"结果已保存到: {save_path}") plt.show() def extract_object_mask(self, pred_seg, target_classes): """ 提取特定类别的mask Args: pred_seg: 分割结果 (H, W) target_classes: 目标类别ID列表 Returns: mask: 二值mask (H, W) """ mask = np.isin(pred_seg, target_classes) return mask.astype(np.uint8) * 255 def extract_food_region(self, image, pred_seg): """ 提取食物区域(基于ADE20K数据集) ADE20K中食物相关类别(部分): - 53: food, 54: table, 55: chair, 56: car, 57: person - 116: plate, 117: banana, 118: apple, 119: sandwich, 120: orange Args: image: 原始图像 (H, W, 3) pred_seg: 分割结果 (H, W) Returns: food_image: 提取的食物区域图像 food_mask: 食物区域mask """ # 食物相关类别(需要根据实际情况调整) food_classes = [53, 116, 117, 118, 119, 120, 121, 122, 123] # 提取mask food_mask = self.extract_object_mask(pred_seg, food_classes) # 应用mask food_image = image.copy() food_image[food_mask == 0] = 255 # 白色背景 return food_image, food_mask def main(): """ 主函数:演示SegFormer-B0的基本用法 """ print("=" * 50) print("SegFormer-B0 快速上手Demo") print("=" * 50) # 1. 创建Demo实例 demo = SegFormerDemo() # 2. 测试图像路径(请替换为您的图像路径) test_image = "test_food3.jpg" # 如果没有测试图像,创建一个示例 if not os.path.exists(test_image): print(f"\n警告: 未找到测试图像 '{test_image}'") print("请将您的食物图片重命名为 'test_food.jpg' 或修改代码中的路径") print("\n提示: 您可以从以下位置获取测试图像:") print(" - dataset目录中的训练图像") print(" - 任何食物图片") return # 3. 进行分割 print(f"\n正在处理图像: {test_image}") pred_seg, original_image = demo.segment_image(test_image) # 4. 显示统计信息 unique_classes = np.unique(pred_seg) print(f"\n检测到的类别数: {len(unique_classes)}") print(f"类别ID: {unique_classes[:20]}...") # 显示前20个 # 5. 可视化结果 print("\n可视化分割结果...") demo.visualize_segmentation( original_image, pred_seg, save_path="segmentation_result.png" ) # 6. 提取食物区域 print("\n提取食物区域...") food_image, food_mask = demo.extract_food_region(original_image, pred_seg) # 可视化食物区域 fig, axes = plt.subplots(1, 3, figsize=(15, 5)) axes[0].imshow(original_image) axes[0].set_title("Original") axes[0].axis("off") axes[1].imshow(food_mask, cmap='gray') axes[1].set_title("Food Mask") axes[1].axis("off") axes[2].imshow(food_image) axes[2].set_title("Extracted Food Region") axes[2].axis("off") plt.tight_layout() plt.savefig("food_extraction_result.png", dpi=150) print("食物提取结果已保存到: food_extraction_result.png") plt.show() print("\n" + "=" * 50) print("Demo运行完成!") print("=" * 50) if __name__ == "__main__": main()