138 lines
5.5 KiB
Python
138 lines
5.5 KiB
Python
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_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()
|