525 lines
17 KiB
Python
525 lines
17 KiB
Python
"""
|
||
模型推理测试脚本
|
||
|
||
功能说明:
|
||
1. 加载训练好的SegFormer模型
|
||
2. 对单张图像进行分割预测
|
||
3. 可视化分割结果
|
||
4. 批量测试并计算指标
|
||
|
||
使用场景:
|
||
- 测试Fine-tune后的模型效果
|
||
- 对比预训练模型和Fine-tune模型
|
||
- 实际应用场景的推理
|
||
"""
|
||
|
||
import os
|
||
import sys
|
||
import torch
|
||
import numpy as np
|
||
from PIL import Image
|
||
import matplotlib.pyplot as plt
|
||
from pathlib import Path
|
||
from typing import Optional, Union
|
||
import torch.nn.functional as F
|
||
from transformers import SegformerForSemanticSegmentation, SegformerImageProcessor
|
||
import albumentations as A
|
||
from albumentations.pytorch import ToTensorV2
|
||
|
||
# 设置matplotlib中文字体
|
||
plt.rcParams['font.sans-serif'] = ['SimHei', 'Microsoft YaHei', 'Arial Unicode MS'] # 用来正常显示中文标签
|
||
plt.rcParams['axes.unicode_minus'] = False # 用来正常显示负号
|
||
|
||
# 添加父目录到路径
|
||
sys.path.append(str(Path(__file__).parent.parent / "training"))
|
||
from config import TrainConfig
|
||
|
||
|
||
class SegFormerInference:
|
||
"""SegFormer推理类"""
|
||
|
||
def __init__(
|
||
self,
|
||
model_path: Optional[str] = None,
|
||
pretrained_model: str = "nvidia/segformer-b0-finetuned-ade-512-512",
|
||
num_classes: int = 2,
|
||
image_size: int = 256, # ⚠️ 重要:必须与训练时一致!
|
||
device: str = "auto"
|
||
):
|
||
"""
|
||
初始化推理类
|
||
|
||
Args:
|
||
model_path: Fine-tune后的模型路径(.pt文件)
|
||
如果为None,则使用预训练模型
|
||
pretrained_model: 预训练模型名称(用于加载processor)
|
||
num_classes: 类别数
|
||
image_size: 输入图像尺寸(必须与训练时一致!)
|
||
device: 设备 ('cpu', 'cuda', 'auto')
|
||
"""
|
||
self.num_classes = num_classes
|
||
self.image_size = image_size
|
||
|
||
# 设置设备
|
||
if device == "auto":
|
||
self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||
else:
|
||
self.device = torch.device(device)
|
||
|
||
print(f"使用设备: {self.device}")
|
||
|
||
# ⚠️ 重要:使用与训练时完全一致的预处理
|
||
# 不再使用 SegformerImageProcessor,而是手动构建预处理pipeline
|
||
print(f"构建预处理Pipeline(与训练时一致)")
|
||
print(f" 图像尺寸: {image_size}×{image_size}")
|
||
print(f" 归一化: ImageNet标准(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])")
|
||
|
||
self.transform = A.Compose([
|
||
A.Resize(image_size, image_size), # ⚠️ 关键:必须resize到训练时的尺寸
|
||
A.Normalize(
|
||
mean=[0.485, 0.456, 0.406], # ImageNet标准均值
|
||
std=[0.229, 0.224, 0.225], # ImageNet标准标准差
|
||
),
|
||
ToTensorV2(),
|
||
])
|
||
|
||
# 加载模型
|
||
if model_path:
|
||
print(f"加载Fine-tune模型: {model_path}")
|
||
self.model = self._load_finetuned_model(model_path)
|
||
else:
|
||
print(f"加载预训练模型: {pretrained_model}")
|
||
self.model = SegformerForSemanticSegmentation.from_pretrained(
|
||
pretrained_model,
|
||
num_labels=num_classes
|
||
)
|
||
|
||
self.model.to(self.device)
|
||
self.model.eval()
|
||
|
||
print("✓ 模型加载完成")
|
||
|
||
def _load_finetuned_model(self, checkpoint_path: str):
|
||
"""
|
||
加载Fine-tune后的模型
|
||
|
||
Args:
|
||
checkpoint_path: checkpoint文件路径
|
||
|
||
Returns:
|
||
model: 加载权重后的模型
|
||
"""
|
||
from transformers import SegformerConfig
|
||
|
||
# 加载checkpoint
|
||
checkpoint = torch.load(checkpoint_path, map_location=self.device, weights_only=False)
|
||
|
||
# 从checkpoint中获取配置
|
||
if 'config' in checkpoint:
|
||
config_dict = checkpoint['config']
|
||
pretrained_model = config_dict.get('pretrained_model', 'nvidia/segformer-b0-finetuned-ade-512-512')
|
||
else:
|
||
pretrained_model = 'nvidia/segformer-b0-finetuned-ade-512-512'
|
||
|
||
# 创建模型配置(不加载预训练权重)
|
||
config = SegformerConfig.from_pretrained(pretrained_model)
|
||
config.num_labels = self.num_classes
|
||
|
||
# 用配置创建空模型
|
||
model = SegformerForSemanticSegmentation(config)
|
||
|
||
# 加载你训练好的权重
|
||
model.load_state_dict(checkpoint['model_state_dict'])
|
||
|
||
# 打印训练信息
|
||
print(f" ✓ 模型权重加载成功")
|
||
if 'epoch' in checkpoint:
|
||
print(f" 训练轮数: {checkpoint['epoch']}")
|
||
if 'metrics' in checkpoint:
|
||
metrics = checkpoint['metrics']
|
||
if 'miou' in metrics:
|
||
print(f" mIoU: {metrics['miou']:.4f}")
|
||
if 'pixel_accuracy' in metrics:
|
||
print(f" 像素准确率: {metrics['pixel_accuracy']:.4f}")
|
||
|
||
return model
|
||
|
||
@torch.no_grad()
|
||
def predict(self, image: Union[str, Image.Image, np.ndarray]) -> np.ndarray:
|
||
"""
|
||
对单张图像进行预测
|
||
|
||
Args:
|
||
image: 输入图像
|
||
- str: 图像文件路径
|
||
- PIL.Image: PIL图像对象
|
||
- np.ndarray: numpy数组 (H, W, 3)
|
||
|
||
Returns:
|
||
pred_mask: 预测的分割mask (H, W),像素值为类别ID
|
||
"""
|
||
# 1. 读取图像并保存原始尺寸
|
||
if isinstance(image, str):
|
||
image = Image.open(image).convert('RGB')
|
||
original_size = image.size # (W, H)
|
||
image = np.array(image)
|
||
elif isinstance(image, np.ndarray):
|
||
original_size = (image.shape[1], image.shape[0]) # (W, H)
|
||
else: # PIL.Image
|
||
original_size = image.size # (W, H)
|
||
image = np.array(image)
|
||
|
||
# 2. 应用与训练时一致的预处理
|
||
# ⚠️ 关键:使用ImageNet标准归一化,与训练时完全一致
|
||
transformed = self.transform(image=image)
|
||
image_tensor = transformed['image'] # (3, H, W)
|
||
|
||
# 3. 添加batch维度
|
||
image_tensor = image_tensor.unsqueeze(0).to(self.device) # (1, 3, H, W)
|
||
|
||
# 4. 模型推理
|
||
outputs = self.model(pixel_values=image_tensor)
|
||
logits = outputs.logits # (1, num_classes, H/4, W/4)
|
||
|
||
# 5. 上采样到原始尺寸
|
||
logits = F.interpolate(
|
||
logits,
|
||
size=(original_size[1], original_size[0]), # (H, W)
|
||
mode='bilinear',
|
||
align_corners=False
|
||
)
|
||
|
||
# 6. 获取预测类别
|
||
pred_mask = logits.argmax(dim=1)[0].cpu().numpy()
|
||
|
||
return pred_mask
|
||
|
||
def visualize_prediction(
|
||
self,
|
||
image_path: str,
|
||
pred_mask: Optional[np.ndarray] = None,
|
||
save_path: Optional[str] = None,
|
||
class_names: list = None
|
||
):
|
||
"""
|
||
可视化预测结果
|
||
|
||
Args:
|
||
image_path: 原始图像路径
|
||
pred_mask: 预测mask(如果为None,则重新预测)
|
||
save_path: 保存路径
|
||
class_names: 类别名称列表
|
||
"""
|
||
# 读取原始图像
|
||
image = Image.open(image_path).convert('RGB')
|
||
image_np = np.array(image)
|
||
|
||
# 预测
|
||
if pred_mask is None:
|
||
pred_mask = self.predict(image)
|
||
|
||
# 默认类别名称
|
||
if class_names is None:
|
||
class_names = ['背景', '食物']
|
||
|
||
# 创建可视化
|
||
fig, axes = plt.subplots(1, 3, figsize=(18, 6))
|
||
|
||
# 1. 原始图像
|
||
axes[0].imshow(image_np)
|
||
axes[0].set_title('原始图像', fontsize=14)
|
||
axes[0].axis('off')
|
||
|
||
# 2. 预测mask(彩色)
|
||
axes[1].imshow(pred_mask, cmap='tab20', vmin=0, vmax=self.num_classes-1)
|
||
axes[1].set_title('预测Mask', fontsize=14)
|
||
axes[1].axis('off')
|
||
|
||
# 3. 叠加显示
|
||
# 创建彩色mask
|
||
colored_mask = np.zeros_like(image_np)
|
||
colors = [
|
||
[0, 0, 0], # 背景:黑色
|
||
[255, 0, 0], # 食物:红色
|
||
]
|
||
for class_id in range(self.num_classes):
|
||
colored_mask[pred_mask == class_id] = colors[class_id]
|
||
|
||
# 叠加
|
||
alpha = 0.5
|
||
overlay = (image_np * (1 - alpha) + colored_mask * alpha).astype(np.uint8)
|
||
axes[2].imshow(overlay)
|
||
axes[2].set_title('叠加显示(红色=食物)', fontsize=14)
|
||
axes[2].axis('off')
|
||
|
||
# 添加统计信息
|
||
total_pixels = pred_mask.size
|
||
food_pixels = np.sum(pred_mask == 1)
|
||
food_ratio = food_pixels / total_pixels * 100
|
||
|
||
plt.suptitle(
|
||
f'分割结果 | 食物区域: {food_ratio:.1f}%',
|
||
fontsize=16,
|
||
y=0.98
|
||
)
|
||
|
||
plt.tight_layout()
|
||
|
||
if save_path:
|
||
plt.savefig(save_path, dpi=150, bbox_inches='tight')
|
||
print(f"✓ 结果已保存: {save_path}")
|
||
|
||
plt.show()
|
||
|
||
def extract_food_region(
|
||
self,
|
||
image_path: str,
|
||
pred_mask: Optional[np.ndarray] = None,
|
||
save_path: Optional[str] = None
|
||
) -> np.ndarray:
|
||
"""
|
||
提取食物区域(去除背景)
|
||
|
||
Args:
|
||
image_path: 原始图像路径
|
||
pred_mask: 预测mask(如果为None,则重新预测)
|
||
save_path: 保存路径
|
||
|
||
Returns:
|
||
food_region: 提取的食物区域图像(背景为白色)
|
||
"""
|
||
# 读取图像
|
||
image = Image.open(image_path).convert('RGB')
|
||
image_np = np.array(image)
|
||
|
||
# 预测
|
||
if pred_mask is None:
|
||
pred_mask = self.predict(image)
|
||
|
||
# 创建二值mask
|
||
food_mask = (pred_mask == 1)
|
||
|
||
# 提取食物区域(白色背景)
|
||
food_region = image_np.copy()
|
||
food_region[~food_mask] = 255
|
||
|
||
# 保存
|
||
if save_path:
|
||
Image.fromarray(food_region).save(save_path)
|
||
print(f"✓ 食物区域已保存: {save_path}")
|
||
|
||
return food_region
|
||
|
||
def batch_test(
|
||
self,
|
||
image_dir: str,
|
||
output_dir: str,
|
||
max_images: int = 10
|
||
):
|
||
"""
|
||
批量测试图像
|
||
|
||
Args:
|
||
image_dir: 图像目录
|
||
output_dir: 输出目录
|
||
max_images: 最大测试图像数
|
||
"""
|
||
image_dir = Path(image_dir)
|
||
output_dir = Path(output_dir)
|
||
output_dir.mkdir(parents=True, exist_ok=True)
|
||
|
||
# 获取所有图像文件
|
||
image_files = []
|
||
for ext in ['*.jpg', '*.jpeg', '*.png']:
|
||
image_files.extend(list(image_dir.glob(ext)))
|
||
|
||
image_files = sorted(image_files)[:max_images]
|
||
|
||
print(f"\n开始批量测试: {len(image_files)} 张图像")
|
||
print("="*60)
|
||
|
||
for idx, image_path in enumerate(image_files, 1):
|
||
print(f"[{idx}/{len(image_files)}] {image_path.name}")
|
||
|
||
# 预测
|
||
pred_mask = self.predict(str(image_path))
|
||
|
||
# 可视化
|
||
save_path = output_dir / f"{image_path.stem}_result.png"
|
||
self.visualize_prediction(
|
||
str(image_path),
|
||
pred_mask=pred_mask,
|
||
save_path=str(save_path)
|
||
)
|
||
|
||
print("="*60)
|
||
print(f"✓ 批量测试完成! 结果保存在: {output_dir}")
|
||
|
||
|
||
def compare_models(
|
||
image_path: str,
|
||
finetuned_model_path: str,
|
||
pretrained_model: str = "nvidia/segformer-b0-finetuned-ade-512-512"
|
||
):
|
||
"""
|
||
对比预训练模型和Fine-tune模型的效果
|
||
|
||
Args:
|
||
image_path: 测试图像路径
|
||
finetuned_model_path: Fine-tune模型路径
|
||
pretrained_model: 预训练模型名称
|
||
"""
|
||
print("\n" + "="*60)
|
||
print("模型对比测试")
|
||
print("="*60)
|
||
|
||
# 加载预训练模型
|
||
print("\n加载预训练模型...")
|
||
pretrained_inference = SegFormerInference(
|
||
model_path=None,
|
||
pretrained_model=pretrained_model,
|
||
num_classes=150, # ADE20K的类别数
|
||
image_size=256 # 与Fine-tune模型保持一致
|
||
)
|
||
|
||
# 加载Fine-tune模型
|
||
print("\n加载Fine-tune模型...")
|
||
finetuned_inference = SegFormerInference(
|
||
model_path=finetuned_model_path,
|
||
pretrained_model=pretrained_model,
|
||
num_classes=2,
|
||
image_size=256 # ⚠️ 必须与训练时一致
|
||
)
|
||
|
||
# 读取图像
|
||
image = Image.open(image_path).convert('RGB')
|
||
image_np = np.array(image)
|
||
|
||
# 预测
|
||
print("\n执行预测...")
|
||
pretrained_mask = pretrained_inference.predict(image)
|
||
finetuned_mask = finetuned_inference.predict(image)
|
||
|
||
# 可视化对比
|
||
fig, axes = plt.subplots(2, 3, figsize=(18, 12))
|
||
|
||
# 预训练模型结果
|
||
axes[0, 0].imshow(image_np)
|
||
axes[0, 0].set_title('原始图像 (预训练模型)', fontsize=12)
|
||
axes[0, 0].axis('off')
|
||
|
||
axes[0, 1].imshow(pretrained_mask, cmap='tab20b')
|
||
axes[0, 1].set_title('预训练模型预测\n(ADE20K, 150类)', fontsize=12)
|
||
axes[0, 1].axis('off')
|
||
|
||
axes[0, 2].imshow(image_np)
|
||
axes[0, 2].imshow(pretrained_mask, alpha=0.5, cmap='tab20b')
|
||
axes[0, 2].set_title('预训练模型叠加', fontsize=12)
|
||
axes[0, 2].axis('off')
|
||
|
||
# Fine-tune模型结果
|
||
axes[1, 0].imshow(image_np)
|
||
axes[1, 0].set_title('原始图像 (Fine-tune模型)', fontsize=12)
|
||
axes[1, 0].axis('off')
|
||
|
||
axes[1, 1].imshow(finetuned_mask, cmap='gray', vmin=0, vmax=1)
|
||
axes[1, 1].set_title('Fine-tune模型预测\n(食物分割, 2类)', fontsize=12)
|
||
axes[1, 1].axis('off')
|
||
|
||
# Fine-tune叠加
|
||
overlay = image_np.copy().astype(float)
|
||
red_mask = np.zeros_like(overlay)
|
||
red_mask[finetuned_mask == 1] = [255, 0, 0]
|
||
overlay = (overlay * 0.6 + red_mask * 0.4).astype(np.uint8)
|
||
axes[1, 2].imshow(overlay)
|
||
axes[1, 2].set_title('Fine-tune模型叠加\n(红色=食物)', fontsize=12)
|
||
axes[1, 2].axis('off')
|
||
|
||
plt.suptitle('模型对比:预训练 vs Fine-tune', fontsize=16, y=0.98)
|
||
plt.tight_layout()
|
||
|
||
save_path = "model_comparison.png"
|
||
plt.savefig(save_path, dpi=150, bbox_inches='tight')
|
||
print(f"\n✓ 对比结果已保存: {save_path}")
|
||
plt.show()
|
||
|
||
|
||
def main():
|
||
"""
|
||
主函数:演示推理用法
|
||
"""
|
||
print("="*60)
|
||
print("SegFormer模型推理测试")
|
||
print("="*60)
|
||
|
||
# ==================== 配置区 ====================
|
||
# TODO: 根据实际情况修改以下配置
|
||
|
||
# Fine-tune模型路径(训练完成后的checkpoint)
|
||
FINETUNED_MODEL_PATH = "d:/MyProjects/PythonProjects/FoodClassifier/SegFormer/outputs/best_model.pt"
|
||
|
||
# 测试图像路径
|
||
TEST_IMAGE_PATH = "d:/MyProjects/PythonProjects/FoodClassifier/SegFormer/test_food4.jpg"
|
||
|
||
# 批量测试目录
|
||
TEST_IMAGE_DIR = "d:/MyProjects/PythonProjects/FoodClassifier/SegFormer/data/segformer_format/images/val"
|
||
OUTPUT_DIR = "d:/MyProjects/PythonProjects/FoodClassifier/SegFormer/inference_results"
|
||
|
||
# ===============================================
|
||
|
||
# 检查模型文件是否存在
|
||
if not os.path.exists(FINETUNED_MODEL_PATH):
|
||
print(f"\n⚠️ Fine-tune模型不存在: {FINETUNED_MODEL_PATH}")
|
||
print(f"\n请先运行训练脚本:")
|
||
print(f" python training/train_minimal.py")
|
||
print(f"\n现在将使用预训练模型进行测试...")
|
||
FINETUNED_MODEL_PATH = None
|
||
|
||
# 创建推理实例
|
||
inference = SegFormerInference(
|
||
model_path=FINETUNED_MODEL_PATH,
|
||
num_classes=2,
|
||
image_size=256 # ⚠️ 必须与训练时一致(见config.py第272行)
|
||
)
|
||
|
||
# 单张图像测试
|
||
if os.path.exists(TEST_IMAGE_PATH):
|
||
print(f"\n测试图像: {TEST_IMAGE_PATH}")
|
||
inference.visualize_prediction(
|
||
TEST_IMAGE_PATH,
|
||
save_path="test_result.png"
|
||
)
|
||
|
||
# 提取食物区域
|
||
inference.extract_food_region(
|
||
TEST_IMAGE_PATH,
|
||
save_path="food_region.png"
|
||
)
|
||
else:
|
||
print(f"\n⚠️ 测试图像不存在: {TEST_IMAGE_PATH}")
|
||
|
||
# 批量测试
|
||
if os.path.exists(TEST_IMAGE_DIR):
|
||
print(f"\n批量测试目录: {TEST_IMAGE_DIR}")
|
||
inference.batch_test(
|
||
image_dir=TEST_IMAGE_DIR,
|
||
output_dir=OUTPUT_DIR,
|
||
max_images=5
|
||
)
|
||
|
||
# 模型对比(如果有Fine-tune模型)
|
||
if FINETUNED_MODEL_PATH and os.path.exists(FINETUNED_MODEL_PATH) and os.path.exists(TEST_IMAGE_PATH):
|
||
print(f"\n执行模型对比...")
|
||
compare_models(
|
||
image_path=TEST_IMAGE_PATH,
|
||
finetuned_model_path=FINETUNED_MODEL_PATH
|
||
)
|
||
|
||
print("\n" + "="*60)
|
||
print("✓ 测试完成!")
|
||
print("="*60)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|