import torch import numpy as np from PIL import Image import matplotlib.pyplot as plt import argparse from src import UNet from src.quality.quality_scorer import QualityScorer # from src.quality.quality_scorer_old import QualityScorer def load_model(model_path, in_channels=3, out_classes=2, device='cpu'): """Загрузка модели""" model = UNet(in_channels=in_channels, out_classes=out_classes) model.load_state_dict(torch.load(model_path, map_location=device)) model.to(device) model.eval() return model def predict_single(model, image_path, device='cpu', input_size=(256, 256)): """Предсказание для одного изображения""" # Загрузка изображения image = Image.open(image_path).convert('RGB') original_image = np.array(image) original_size = image.size # (width, height) # Подготовка для модели (ресайз) image_resized = image.resize(input_size) image_array = np.array(image_resized, dtype=np.float32) / 255.0 image_array = image_array.transpose(2, 0, 1) image_tensor = torch.FloatTensor(image_array).unsqueeze(0).to(device) # Инференс with torch.no_grad(): output = model(image_tensor) pred = torch.softmax(output, dim=1) mask_resized = pred.argmax(dim=1).squeeze().cpu().numpy() # Ресайз маски обратно к оригинальному размеру # Используем NEAREST интерполяцию для сохранения бинарности mask_pil = Image.fromarray(mask_resized.astype(np.uint8)) mask_pil = mask_pil.resize(original_size, resample=Image.NEAREST) mask = np.array(mask_pil) return original_image, mask def visualize_prediction(image, mask, quality_scores, save_path=None): """Визуализация результатов""" fig, axes = plt.subplots(1, 3, figsize=(15, 5)) # 1. Оригинальное изображение axes[0].imshow(image) axes[0].set_title('Оригинал') axes[0].axis('off') # 2. Сегментация (маска наложена на оригинал) axes[1].imshow(image) axes[1].imshow(mask, cmap='jet', alpha=0.5) axes[1].set_title('Сегментация') axes[1].axis('off') # 3. Оценка качества quality_text = f"Качество: {quality_scores['overall_quality']}\n" quality_text += f"Уверенность: {quality_scores['confidence']:.2f}\n" quality_text += f"Нарушений: {len(quality_scores['issues'])}\n\n" # Добавляем детали if quality_scores['issues']: quality_text += "Проблемы:\n" for issue in quality_scores['issues']: quality_text += f" • {issue['type']}: {issue['details'][:30]}...\n" else: quality_text += "✅ Нарушений не обнаружено" axes[2].text(0.05, 0.95, quality_text, fontsize=10, verticalalignment='top', transform=axes[2].transAxes, bbox=dict(boxstyle="round", facecolor="white", alpha=0.8)) axes[2].set_title('Оценка качества') axes[2].axis('off') plt.rcParams['font.family'] = 'sans-serif' plt.rcParams['font.sans-serif'] = ['Apple Color Emoji', 'Segoe UI Emoji', 'Noto Color Emoji', 'DejaVu Sans'] plt.tight_layout() if save_path: plt.savefig(save_path, dpi=150, bbox_inches='tight') print(f"✅ Визуализация сохранена: {save_path}") plt.show() def main(): parser = argparse.ArgumentParser(description='Инференс модели сегментации') parser.add_argument('--image', type=str, default="example.jpg", required=False, help='Путь к изображению') parser.add_argument('--model', type=str, default='models/unet_cats_dogs.pth', help='Путь к модели') parser.add_argument('--device', type=str, default='cpu', choices=['cpu', 'mps', 'cuda'], help='Устройство') parser.add_argument('--save', type=str, default=None, help='Сохранить результат') args = parser.parse_args() # Проверка устройства if args.device == 'mps' and not torch.backends.mps.is_available(): print("⚠️ MPS не доступен, используем CPU") args.device = 'cpu' print(f"🚀 Загрузка модели с {args.device}...") model = load_model(args.model, device=args.device) print(f"📸 Анализ изображения: {args.image}") image, mask = predict_single(model, args.image, device=args.device) print(f" Изображение: {image.shape}") print(f" Маска: {mask.shape}") print("🔍 Оценка качества...") scorer = QualityScorer() quality_scores = scorer.evaluate(image=image, segmentation=mask) print(f"\n📊 Результаты:") print(f" Качество: {quality_scores['overall_quality']}") print(f" Уверенность: {quality_scores['confidence']:.2f}") print(f" Количество проблем: {len(quality_scores['issues'])}") if quality_scores['issues']: print("\n Проблемы:") for issue in quality_scores['issues']: print(f" • {issue['type']}: {issue['details'][:50]}...") # Визуализация save_path = args.save or 'prediction_result.png' visualize_prediction(image, mask, quality_scores, save_path) if __name__ == "__main__": main()