bone_2026/inference.py

141 lines
5.5 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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()