commit a5ada071dc4a1158fb8e4a887aae6dfb48fe2f23 Author: denis Date: Sat Aug 15 23:20:59 2026 +0300 develop diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..6fd0b2b --- /dev/null +++ b/.dockerignore @@ -0,0 +1,30 @@ +# .dockerignore +__pycache__ +*.pyc +*.pyo +*.pyd +.Python +*.so +*.egg +*.egg-info +dist +build +.venv +venv +env +.env +.git +.gitignore +*.md +.DS_Store +*.log +*.pkl +*.h5 +*.t7 +data/ +datasets/ +.DS_Store +.idea/ +.vscode/ +*.swp +*.swo \ No newline at end of file diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..7356404 --- /dev/null +++ b/Dockerfile @@ -0,0 +1,36 @@ +# Dockerfile +FROM python:3.10-slim + +# Устанавливаем рабочую директорию +WORKDIR /app + +# Устанавливаем системные зависимости для OpenCV и PIL +RUN apt-get update && apt-get install -y \ + libgl1-mesa-glx \ + libglib2.0-0 \ + libsm6 \ + libxext6 \ + libxrender-dev \ + libgomp1 \ + wget \ + && rm -rf /var/lib/apt/lists/* + +# Копируем requirements.txt и устанавливаем зависимости +COPY requirements.txt . +RUN pip install --no-cache-dir -r requirements.txt --extra-index-url https://download.pytorch.org/whl/cpu + +# Копируем весь проект +COPY . . + +# Создаем папку для моделей (если её нет) +RUN mkdir -p models + +# Открываем порт +EXPOSE 8000 + +# Переменные окружения +ENV PYTHONUNBUFFERED=1 +ENV PYTHONPATH=/app + +# Запуск сервера +CMD ["python", "run.py"] \ No newline at end of file diff --git a/README.md b/README.md new file mode 100644 index 0000000..5a9ef6a --- /dev/null +++ b/README.md @@ -0,0 +1,256 @@ +# 🦴 Bone Quality Assessment + +[![Python](https://img.shields.io/badge/Python-3.10-blue.svg)](https://www.python.org/) +[![FastAPI](https://img.shields.io/badge/FastAPI-0.104-green.svg)](https://fastapi.tiangolo.com/) +[![PyTorch](https://img.shields.io/badge/PyTorch-2.1-red.svg)](https://pytorch.org/) +[![Docker](https://img.shields.io/badge/Docker-Ready-blue.svg)](https://www.docker.com/) + +## 📋 Описание + +Сервис искусственного интеллекта для автоматизированной оценки качества денситометрических изображений и их разметки. Система получает на вход рентгеновское денситометрическое исследование в формате DICOM, и оценивает качество выполнения исследования по стандартным критериям, а также корректность разметки анатомических структур на изображениях. + +### Основные возможности + +- 🖼️ **Анализ изображений** — загрузка и обработка медицинских изображений +- 🧠 **Сегментация объектов** — выделение анатомических структур (позвонки, кости) +- 📊 **Оценка качества** — проверка по 3 критериям: + - Артефакты (движение, шум, размытость) + - Позиционирование (правильное расположение объекта) + - Контрастность (качество изображения) +- 🔍 **Детекция нарушений** — определение типа нарушения для некачественных исследований +- 🌐 **Web-интерфейс** — удобная загрузка и визуализация результатов +- 📡 **REST API** — интеграция с внешними системами +- 📈 **Визуализация маски** — отображение сегментации на изображении + +## 🏗️ Архитектура + +![!img](public/static/arch.png) + + +## 🚀 Быстрый старт + +### Требования + +- Python 3.10+ +- PyTorch 2.1+ +- Docker (опционально) + +### Локальная установка + +```bash +# 1. Клонирование репозитория +git clone https://github.com/yourusername/bone-quality-assessment.git +cd bone-quality-assessment + +# 2. Создание виртуального окружения +python -m venv venv +source venv/bin/activate # Linux/Mac +# или +venv\Scripts\activate # Windows + +# 3. Установка зависимостей +pip install -r requirements.txt + +# 4. Загрузка обученной модели (опционально) +# Поместите модель в папку models/unet_cats_dogs.pth + +# 5. Запуск сервера +python run.py +``` + +Docker + +```bash +# 1. Сборка образа +docker build -t bone-quality-api . + +# 2. Запуск контейнера +docker run -p 8000:8000 bone-quality-api + +# 3. Или используя docker-compose +docker-compose up -d +``` + + +📡 API Endpoints + +| Метод | Эндпоинт | Описание | +|--------|-------------------|--------------------------| +| GET | / | Главная страница +| GET | /docs | Swagger UI документация +| GET | /redoc | ReDoc документация +| GET | /api/v1/health | Проверка статуса сервиса +| POST | /api/v1/analyze | Анализ изображения + + +Пример запроса +```bash +curl -X POST "http://localhost:8000/api/v1/analyze" \ + -H "accept: application/json" \ + -H "Content-Type: multipart/form-data" \ + -F "file=@/path/to/image.jpg" +``` + +Пример ответа + +```json +{ + "overall_quality": "GOOD", + "severity": "LOW", + "issues": [], + "confidence": 0.9, + "metrics": { + "artifact": { + "artifact": false, + "num_objects": 1, + "edge_energy": 0.234, + "object_size": 0.123 + }, + "position": { + "position": [0.45, 0.52], + "valid": true, + "deviation": 0.032 + }, + "contrast": { + "valid": true, + "contrast": 0.456 + } + }, + "mask": "base64_encoded_mask_image" +} +``` + + +📊 Интерфейс +Веб-интерфейс доступен по адресу http://localhost:8000/: + + - 📤 Drag-and-drop загрузка изображений + - 🔍 Автоматический анализ + - 🎯 Визуализация маски сегментации + - 📈 Детальные метрики качества + - 🏷️ Подробный отчет о нарушениях + + +🛠️ Технологии + +|Компонент | Технология +|-|-| +|Бэкенд | Python 3.10, FastAPI, Uvicorn +|ML | PyTorch, NumPy, SciPy +|Обработка изображений | PIL, OpenCV +|Визуализация | HTML5, CSS3, Canvas API +|Контейнеризация | Docker, Docker Compose +|Документация | Swagger UI, ReDoc + + +📁 Структура проекта +```text +bone-quality-assessment/ +├── src/ +│ ├── api/ +│ │ ├── endpoints.py # FastAPI эндпоинты +│ │ └── static/ +│ │ └── index.html # Web-интерфейс +│ ├── models/ +│ │ └── unet.py # U-Net архитектура +│ └── quality/ +│ └── quality_scorer.py # Оценка качества +├── models/ +│ └── unet_cats_dogs.pth # Обученная модель +├── Dockerfile +├── docker-compose.yml +├── requirements.txt +├── run.py +└── README.md +``` + +🧪 Тестирование +```bash +# Запуск тестов (если есть) +pytest tests/ +``` + +# Проверка API + +```bash +curl http://localhost:8000/api/v1/health +``` + +## 📊 Метрики качества + +Артефакты + +*Описание:* Обнаружение шума, размытости, фрагментации + +Проверка: + + - Энергия границ (edge_energy) + - Количество объектов (num_objects) + - Размер объекта (object_size) + +### Позиционирование + +*Описание:* Проверка правильности расположения объекта + +Проверка: + + - Центр масс объекта + - Отклонение от центра изображения + - Минимальный размер объекта + +### Контраст + +*Описание:* Оценка качества изображения + +Проверка: + + - Контраст внутри объекта + - Отношение средних (объект/фон) + - Интенсивность пикселей + +## 🚀 Деплой +На сервер +```bash +# Копирование на сервер +scp -r ./bone-quality-assessment user@server:/var/www/ +``` + +# Запуск в фоновом режиме +nohup python run.py > logs/out.log 2>&1 & + +Использование с Nginx +```nginx +location /api/ { + proxy_pass http://localhost:8000; + proxy_set_header Host $host; + proxy_set_header X-Real-IP $remote_addr; +} +``` + +## 🤝 Вклад в проект + - Fork репозитория + - Создайте ветку для вашей фичи (git checkout -b feature/amazing-feature) + - Commit изменений (git commit -m 'Add amazing feature') + - Push в ветку (git push origin feature/amazing-feature) + - Откройте Pull Request + + +**📄 Лицензия** + +MIT License + +👥 Команда +Грачев Денис — Разработка - GitHub + +### 🙏 Благодарности + +***Oxford-IIIT Pet Dataset*** для обучения модели + +***Сообществу PyTorch и FastAPI*** + +📞 Контакты + - 📧 Email: your.email@example.com + - 🐦 Telegram: @oxydencher + - 🐙 GitHub: gdg6 + +
Built with ❤️ for the Bone Quality Assessment Hackathon
diff --git a/example.jpg b/example.jpg new file mode 100755 index 0000000..b4acd1a Binary files /dev/null and b/example.jpg differ diff --git a/example2.jpg b/example2.jpg new file mode 100644 index 0000000..cb92134 Binary files /dev/null and b/example2.jpg differ diff --git a/inference.py b/inference.py new file mode 100644 index 0000000..1a31822 --- /dev/null +++ b/inference.py @@ -0,0 +1,137 @@ +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() diff --git a/main.py b/main.py new file mode 100644 index 0000000..8f8ca2f --- /dev/null +++ b/main.py @@ -0,0 +1,6 @@ +# Press the green button in the gutter to run the script. +from pydicom_custom import main + +if __name__ == '__main__': + main() +# See PyCharm help at https://www.jetbrains.com/help/pycharm/ diff --git a/prediction_result.png b/prediction_result.png new file mode 100644 index 0000000..1ba652c Binary files /dev/null and b/prediction_result.png differ diff --git a/public/static/arch.png b/public/static/arch.png new file mode 100644 index 0000000..fcbe4eb Binary files /dev/null and b/public/static/arch.png differ diff --git a/pydicom_custom.py b/pydicom_custom.py new file mode 100644 index 0000000..5a64808 --- /dev/null +++ b/pydicom_custom.py @@ -0,0 +1,19 @@ +import pydicom +import matplotlib.pyplot as plt +import numpy as np +def main(): + # Читаем DICOM-файл + ds = pydicom.dcmread("./datasets/MRBRAIN.DCM") + + # Получаем пиксельные данные в виде NumPy-массива + image_array = ds.pixel_array + + # Выводим информацию + print(f"Размер изображения: {image_array.shape}") + print(f"Пациент: {ds.PatientName}") + print(f"Модальность: {ds.Modality}") + + # Отображаем изображение + plt.imshow(image_array, cmap='gray') + plt.axis('off') + plt.show() \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..bbafbbf --- /dev/null +++ b/requirements.txt @@ -0,0 +1,47 @@ +annotated-doc==0.0.5 +annotated-types==0.7.0 +anyio==4.12.1 +click==8.1.8 +contourpy==1.3.0 +cycler==0.12.1 +exceptiongroup==1.3.1 +fastapi==0.128.8 +filelock==3.19.1 +fonttools==4.60.2 +fsspec==2025.10.0 +h11==0.16.0 +idna==3.18 +importlib_resources==6.5.2 +Jinja2==3.1.6 +kiwisolver==1.4.7 +MarkupSafe==3.0.3 +matplotlib==3.9.4 +monai==1.5.2 +mpmath==1.3.0 +networkx==3.2.1 +nibabel==5.3.3 +numpy==2.0.2 +opencv-python-headless==5.0.0.93 +packaging==26.3 +pandas==2.3.3 +pillow==11.3.0 +pydantic==2.13.4 +pydantic_core==2.46.4 +pydicom==2.4.4 +pyparsing==3.3.2 +python-dateutil==2.9.0.post0 +python-multipart==0.0.20 +pytz==2026.3.post1 +scipy==1.13.1 +six==1.17.0 +starlette==0.49.3 +sympy==1.14.0 +torch==2.8.0 +torchvision==0.23.0 +TotalSegmentator==2.18.0 +tqdm==4.70.0 +typing-inspection==0.4.2 +typing_extensions==4.16.0 +tzdata==2026.3 +uvicorn==0.39.0 +zipp==3.23.1 diff --git a/src/__init__.py b/src/__init__.py new file mode 100644 index 0000000..9954658 --- /dev/null +++ b/src/__init__.py @@ -0,0 +1,31 @@ +""" +Bone Quality Assessment - главный пакет +Автоматизированная оценка качества денситометрических исследований +""" + +# Управление версией +__version__ = "1.0.0" +__author__ = "Your Team" + +# Импорт основных компонентов для удобного доступа +from src.config import Config +from src.data_loader import FlexibleDataset +from src.model import SegmentationEngine, UNet +from src.quality import QualityScorer +from src.api import app + +# Что импортируется при "from src import *" +__all__ = [ + "Config", + "FlexibleDataset", + "UNet", + "SegmentationEngine", + "QualityScorer", + "app", + "__version__", +] + +# Можно добавить инициализацию логгера +import logging + +logging.getLogger(__name__).addHandler(logging.NullHandler()) \ No newline at end of file diff --git a/src/api/__init__.py b/src/api/__init__.py new file mode 100644 index 0000000..7aa7f09 --- /dev/null +++ b/src/api/__init__.py @@ -0,0 +1,10 @@ +""" +REST API для сервиса оценки качества +""" + +from src.api.endpoints import app + +# Можно добавить middleware или настройки +__all__ = [ + "app", +] \ No newline at end of file diff --git a/src/api/endpoints.py b/src/api/endpoints.py new file mode 100644 index 0000000..b847d24 --- /dev/null +++ b/src/api/endpoints.py @@ -0,0 +1,267 @@ +from fastapi import FastAPI, File, UploadFile, APIRouter +from fastapi.responses import JSONResponse, FileResponse +from fastapi.staticfiles import StaticFiles +from pydantic import BaseModel +import numpy as np +from PIL import Image +import io +import torch +from pathlib import Path +from typing import List, Optional, Dict, Any +import base64 + +from src import UNet +from src.quality.quality_scorer import QualityScorer, convert_to_serializable + +# Создаем приложение +app = FastAPI( + title="Bone Quality Assessment API", + description="API для оценки качества медицинских изображений", + version="1.0.0", + docs_url="/docs", + redoc_url="/redoc", + openapi_url="/openapi.json" +) + +# Подключаем статические файлы +static_dir = Path(__file__).parent / "static" +if static_dir.exists(): + app.mount("/static", StaticFiles(directory=str(static_dir)), name="static") + print(f"✅ Статика подключена: {static_dir}") +else: + print(f"⚠️ Папка static не найдена: {static_dir}") + +# Создаем роутер +router = APIRouter(prefix="/api/v1", tags=["quality"]) + +# Глобальные переменные +model = None +scorer = None +device = None + + +def get_device(): + """Определение устройства""" + if torch.backends.mps.is_available(): + return torch.device("mps") + elif torch.cuda.is_available(): + return torch.device("cuda") + else: + return torch.device("cpu") + + +def load_models(): + """Ленивая загрузка моделей""" + global model, scorer, device + + if model is None: + device = get_device() + print(f"🔧 Используем устройство: {device}") + + # Загрузка модели + model = UNet(in_channels=3, out_classes=2) + model_path = Path("../../models/unet_cats_dogs.pth") + + if model_path.exists(): + print(f"📦 Загрузка модели из {model_path}") + state_dict = torch.load(model_path, map_location=device) + model.load_state_dict(state_dict) + model.to(device) + model.eval() + print(f"✅ Модель загружена на {device}") + else: + print(f"⚠️ Модель не найдена: {model_path}") + print(" Используем необученную модель (будет работать плохо)") + model.to(device) + + if scorer is None: + scorer = QualityScorer() + print("✅ QualityScorer инициализирован") + + return model, scorer, device + + +class QualityResponse(BaseModel): + overall_quality: str + severity: str + issues: List[Dict[str, Any]] + confidence: float + metrics: Dict[str, Any] + position_validation: Optional[Dict[str, Any]] = None + artifact_validation: Optional[Dict[str, Any]] = None + mask: Optional[str] = None # base64 encoded mask image + + +class HealthResponse(BaseModel): + status: str + model_loaded: bool + device: str + scorer_loaded: bool + model_path_exists: bool + + +@router.post("/analyze", response_model=QualityResponse) +async def analyze_image(file: UploadFile = File(...)): + """Анализ изображения на качество""" + try: + # Загрузка модели + model, scorer, device = load_models() + + # Чтение изображения + image_data = await file.read() + print(f"📸 Загружено {len(image_data)} байт") + + # Конвертация в PIL Image + image = Image.open(io.BytesIO(image_data)).convert('RGB') + print(f" Изображение: {image.size}, {image.mode}") + + # Сохраняем оригинал для визуализации + original_image = np.array(image) + original_size = image.size # (width, height) + + # Подготовка для модели + input_size = (256, 256) + 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) + probs = torch.softmax(output, dim=1) + pred_mask = probs.argmax(dim=1).squeeze().cpu().numpy() + + # Если маска пустая — пробуем порог + if pred_mask.sum() == 0: + print("⚠️ Маска пустая! Пробуем пороговую обработку...") + prob_class1 = probs[0, 1].cpu().numpy() + for threshold in [0.2, 0.15, 0.1]: + temp_mask = (prob_class1 > threshold).astype(np.int64) + if temp_mask.sum() > 0: + print(f" ✅ Найден объект при пороге {threshold}") + pred_mask = temp_mask + break + + # Ресайз маски к оригинальному размеру + mask_pil = Image.fromarray(pred_mask.astype(np.uint8)) + mask_pil = mask_pil.resize(original_size, resample=Image.NEAREST) + mask = np.array(mask_pil) + + print(f" Финальная маска: {mask.shape}, сумма = {mask.sum()}") + + # Конвертируем маску в base64 для передачи на фронтенд + mask_image = Image.fromarray((mask * 255).astype(np.uint8)) + buffer = io.BytesIO() + mask_image.save(buffer, format='PNG') + mask_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8') + + # Оценка качества + quality_scores = scorer.evaluate(image=original_image, segmentation=mask) + + print(f"\n📊 Результаты:") + print(f" Качество: {quality_scores['overall_quality']}") + print(f" Уверенность: {quality_scores['confidence']:.2f}") + print(f" Количество проблем: {len(quality_scores['issues'])}") + + # Конвертируем все numpy типы + quality_scores = convert_to_serializable(quality_scores) + + print(f" Качество (NumPy): {quality_scores['overall_quality']}") + + # Формирование ответа + return QualityResponse( + overall_quality=quality_scores["overall_quality"], + severity=quality_scores.get("severity", "LOW"), + issues=quality_scores.get("issues", []), + confidence=quality_scores.get("confidence", 0.5), + metrics=quality_scores.get("metrics", {}), + position_validation=quality_scores.get("metrics", {}).get("position", {}), + artifact_validation=quality_scores.get("metrics", {}).get("artifact", {}), + mask=mask_base64 + ) + + except Exception as e: + import traceback + traceback.print_exc() + return QualityResponse( + overall_quality="ERROR", + severity="CRITICAL", + issues=[{"type": "PROCESSING_ERROR", "details": str(e), "severity": "critical"}], + confidence=0.0, + metrics={}, + position_validation={"error": str(e)}, + artifact_validation={"error": str(e)}, + mask=None + ) + + +@router.get("/health", response_model=HealthResponse) +async def health_check(): + """Проверка статуса сервиса""" + model, scorer, device = load_models() + model_path = Path("../../models/unet_cats_dogs.pth") + return HealthResponse( + status="ok", + model_loaded=model is not None, + device=str(device), + scorer_loaded=scorer is not None, + model_path_exists=model_path.exists() + ) + + +# Добавляем роутер в приложение +app.include_router(router) + + +# Корневой эндпоинт +@app.get("/") +async def root(): + """Главная страница""" + index_path = static_dir / "index.html" + if index_path.exists(): + return FileResponse(str(index_path)) + else: + return { + "service": "Bone Quality Assessment API", + "version": "1.0.0", + "docs": "/docs", + "health": "/api/v1/health", + "analyze": "/api/v1/analyze (POST)" + } + + +# Обработчик 404 +@app.exception_handler(404) +async def not_found_handler(request, exc): + return JSONResponse( + status_code=404, + content={ + "error": "Not Found", + "message": f"Endpoint {request.url.path} not found", + "available_endpoints": [ + "/", + "/docs", + "/redoc", + "/openapi.json", + "/api/v1/health", + "/api/v1/analyze (POST)", + "/static/" + ] + } + ) + + +if __name__ == "__main__": + import uvicorn + + print("🚀 Запуск Bone Quality Assessment API") + print(f"📍 Статика: {static_dir}") + print("🌐 http://localhost:8000") + uvicorn.run( + "src.api.endpoints:app", + host="0.0.0.0", + port=8000, + reload=True, + log_level="info" + ) \ No newline at end of file diff --git a/src/api/static/index.html b/src/api/static/index.html new file mode 100644 index 0000000..877fb6e --- /dev/null +++ b/src/api/static/index.html @@ -0,0 +1,875 @@ + + + + + + Bone Quality Assessment - Анализ качества костной ткани + + + +
+

🦴 Bone Quality Assessment

+

Загрузите медицинское изображение для анализа качества костной ткани

+ + +
+ 📤 +

Перетащите изображение сюда

+

или кликните для выбора файла

+

Поддерживаются: PNG, JPG, JPEG

+ +
+ + +
+ Preview +
+
+ + • + +
+
+ + + + +
+ + +
+

🧩 Визуализация сегментации

+
+ + + +
+
+ Mask Visualization + +
+
+ 🧊 Объектов: - + 📏 Размер: - + 🎯 Уверенность: - +
+
+ + +
+
+
Общее качество
+
+
+ +
+
Серьезность проблемы
+
+
+ +
+
Выявленные проблемы
+
+
+ +
+
+
📍 Позиционирование
+
+
+
+
+
🔬 Артефакты
+
+
+
+
+ +
+
Уверенность модели
+
+
+
+
+
+ + +
+
+ + + + \ No newline at end of file diff --git a/src/config.py b/src/config.py new file mode 100644 index 0000000..e7c4808 --- /dev/null +++ b/src/config.py @@ -0,0 +1,9 @@ +class Config: + # На хакатоне просто меняешь эти пути! + SEGMENTATION_MODEL = "unet_cats_dogs.pth" # → "totalsegmentator" + DATA_PATH = "data/cats_dogs/" # → "data/dicom/" + INPUT_SIZE = (512, 512) + NUM_CLASSES = 2 # фон + объект + + # Абстрактный интерфейс + USE_TOTAL_SEGMENTATOR = False # → True на хакатоне \ No newline at end of file diff --git a/src/data_loader.py b/src/data_loader.py new file mode 100644 index 0000000..fb47004 --- /dev/null +++ b/src/data_loader.py @@ -0,0 +1,49 @@ +# src/data_loader.py - легко заменить на DICOM +import torch +from torch.utils.data import Dataset, DataLoader +from torchvision import transforms +import numpy as np +from PIL import Image +import os + + +class FlexibleDataset(Dataset): + """Датасет, который можно адаптировать под любые данные""" + + def __init__(self, data_dir, is_train=True, input_size=(512, 512)): + self.data_dir = data_dir + self.input_size = input_size + self.is_train = is_train + + # Собираем все изображения + self.images = [] + self.masks = [] + + # Ищем файлы (эта логика будет меняться под DICOM) + for f in os.listdir(os.path.join(data_dir, 'images')): + if f.endswith(('.jpg', '.png')): + self.images.append(os.path.join(data_dir, 'images', f)) + mask_path = os.path.join(data_dir, 'masks', f.replace('.jpg', '.png')) + if os.path.exists(mask_path): + self.masks.append(mask_path) + + def __len__(self): + return len(self.images) + + def __getitem__(self, idx): + # Загрузка изображения + image = Image.open(self.images[idx]).convert('RGB') + image = image.resize(self.input_size) + image = np.array(image).transpose(2, 0, 1) / 255.0 + + # Загрузка маски (если есть) + if idx < len(self.masks): + mask = Image.open(self.masks[idx]) + mask = mask.resize(self.input_size) + mask = np.array(mask) + # Бинаризация (0 - фон, 1 - объект) + mask = (mask > 128).astype(np.int64) + else: + mask = np.zeros(self.input_size, dtype=np.int64) + + return torch.FloatTensor(image), torch.LongTensor(mask) \ No newline at end of file diff --git a/src/dataloaders/__init__.py b/src/dataloaders/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/src/dataloaders/pet_dataset.py b/src/dataloaders/pet_dataset.py new file mode 100644 index 0000000..9b98288 --- /dev/null +++ b/src/dataloaders/pet_dataset.py @@ -0,0 +1,141 @@ +import os +import torch +from torch.utils.data import Dataset, DataLoader +import numpy as np +from PIL import Image +import matplotlib.pyplot as plt +from pathlib import Path + + +class PetDataset(Dataset): + """Датасет для котиков и собачек из Oxford-IIIT Pet""" + + def __init__(self, data_root, split='trainval', input_size=(512, 512), + binary_mask=True): + """ + Args: + data_root: путь к корневой папке с датасетом + split: 'trainval' или 'test' + input_size: размер для ресайза (H, W) + binary_mask: True -> бинарная маска (фон/объект), + False -> 3 класса (фон/объект/непонятно) + """ + self.data_root = Path(data_root) + self.input_size = input_size + self.binary_mask = binary_mask + + # Пути к папкам + self.images_dir = self.data_root / 'images' + self.trimaps_dir = self.data_root / 'annotations' / 'trimaps' + + # Читаем список файлов для сплита + split_file = self.data_root / 'annotations' / f'{split}.txt' + with open(split_file, 'r') as f: + self.image_names = [line.strip().split()[0] for line in f.readlines()] + + print(f"Загружено {len(self.image_names)} изображений для {split}") + + def __len__(self): + return len(self.image_names) + + def __getitem__(self, idx): + # Получаем имя файла без расширения + img_name = self.image_names[idx] + + # Загружаем изображение (RGB) + img_path = self.images_dir / f'{img_name}.jpg' + image = Image.open(img_path).convert('RGB') + image = image.resize(self.input_size) + image = np.array(image, dtype=np.float32) / 255.0 # Нормализация 0-1 + image = image.transpose(2, 0, 1) # HWC -> CHW + + # Загружаем маску (trimap) + mask_path = self.trimaps_dir / f'{img_name}.png' + mask = Image.open(mask_path) + mask = mask.resize(self.input_size, resample=Image.NEAREST) + mask = np.array(mask, dtype=np.int64) + + if self.binary_mask: + # Бинарная маска: 1 -> объект (все что не фон) + # В оригинале: 1=объект, 2=фон, 3=непонятно + mask = (mask != 2).astype(np.int64) # 1 если не фон, иначе 0 + + return torch.FloatTensor(image), torch.LongTensor(mask) + + def visualize(self, idx): + """Визуализация примера""" + image, mask = self[idx] + + # Переводим обратно в numpy для отображения + image_np = image.numpy().transpose(1, 2, 0) + + fig, axes = plt.subplots(1, 2, figsize=(12, 6)) + + axes[0].imshow(image_np) + axes[0].set_title('Изображение') + axes[0].axis('off') + + axes[1].imshow(mask.numpy(), cmap='tab10', alpha=0.7) + axes[1].set_title('Маска (сегментация)') + axes[1].axis('off') + + plt.show() + + +def create_dataloaders(data_root, batch_size=8, input_size=(512, 512)): + """Создает train и val загрузчики""" + + # Обучающая выборка + train_dataset = PetDataset( + data_root=data_root, + split='trainval', + input_size=input_size, + binary_mask=True + ) + + # Тестовая выборка + val_dataset = PetDataset( + data_root=data_root, + split='test', + input_size=input_size, + binary_mask=True + ) + + train_loader = DataLoader( + train_dataset, + batch_size=batch_size, + shuffle=True, + num_workers=4, + pin_memory=True + ) + + val_loader = DataLoader( + val_dataset, + batch_size=batch_size, + shuffle=False, + num_workers=4, + pin_memory=True + ) + + return train_loader, val_loader + + +if __name__ == "__main__": + # Тестируем даталоадер + data_root = "/Users/denis/workspace/bone_2026/datasets/oxford-iiit-pet" # путь к твоей папке + + train_loader, val_loader = create_dataloaders( + data_root=data_root, + batch_size=4, + input_size=(256, 256) # маленький размер для теста + ) + + # Проверяем один батч + images, masks = next(iter(train_loader)) + print(f"Batch images shape: {images.shape}") # [B, 3, H, W] + print(f"Batch masks shape: {masks.shape}") # [B, H, W] + print(f"Unique mask values: {torch.unique(masks)}") # Должно быть [0, 1] + + # Визуализируем первый пример + dataset = PetDataset(data_root=data_root, split='trainval') + dataset.visualize(0) \ No newline at end of file diff --git a/src/model/__init__.py b/src/model/__init__.py new file mode 100644 index 0000000..7179c4b --- /dev/null +++ b/src/model/__init__.py @@ -0,0 +1,41 @@ +""" +Модели для сегментации +""" + +__all__ = [ + # U-Net + "UNet", + "DoubleConv", + "DownBlock", + "UpBlock", + + "SegmentationEngine" +] + +from src.model.segmentator import SegmentationEngine +from src.model.unet import DoubleConv, UpBlock, DownBlock, UNet + +MODEL_REGISTRY = { + "unet": UNet, + "totalsegmentator": "TotalSegmentator (external)" +} + + +def get_model(name, **kwargs): + """ + Фабричный метод для получения модели по имени + + Args: + name: Имя модели ("unet", "totalsegmentator") + **kwargs: Параметры для модели + + Returns: + Модель для сегментации + """ + if name == "unet": + return UNet(**kwargs) + # elif name == "totalsegmentator": + # from src.models.segmentator import TotalSegmentatorWrapper + # return TotalSegmentatorWrapper(**kwargs) + else: + raise ValueError(f"Unknown model: {name}") \ No newline at end of file diff --git a/src/model/segmentator.py b/src/model/segmentator.py new file mode 100644 index 0000000..83a364d --- /dev/null +++ b/src/model/segmentator.py @@ -0,0 +1,198 @@ +import torch +import numpy as np +import cv2 +from typing import Optional +import os + +from src.model.unet import UNet + + +class SegmentationEngine: + def __init__(self, use_totalsegmentator=False, model_path: Optional[str] = None, device: str = 'cpu'): + """ + Инициализация движка сегментации + + Args: + use_totalsegmentator: использовать TotalSegmentator (если False - UNet) + model_path: путь к файлу модели UNet (.pth) + device: 'cpu' или 'cuda' + """ + self.use_totalsegmentator = use_totalsegmentator + self.device = device if torch.cuda.is_available() else 'cpu' + + # Определяем путь к модели по умолчанию + if model_path is None: + # Ищем модель в стандартных папках + possible_paths = [ + "models/unet_bone.pth", + "../models/unet_bone.pth", + "../../models/unet_bone.pth", + os.path.expanduser("~/models/unet_bone.pth") + ] + for path in possible_paths: + if os.path.exists(path): + model_path = path + break + + if use_totalsegmentator: + self.model = self._load_totalsegmentator() + else: + self.model = self._load_unet(model_path) + + print(f"✅ SegmentationEngine инициализирован на {self.device}") + print(f" Модель: {'TotalSegmentator' if use_totalsegmentator else 'UNet'}") + if model_path: + print(f" Путь: {model_path}") + + def segment(self, image: np.ndarray) -> np.ndarray: + """ + Сегментация изображения + + Args: + image: изображение [H, W] или [H, W, C] + + Returns: + Бинарная маска сегментации [H, W] + """ + if self.use_totalsegmentator: + return self._segment_with_totalsegmentator(image) + else: + return self._segment_with_unet(image) + + def _load_unet(self, model_path: Optional[str] = None) -> torch.nn.Module: + """Загрузка UNet модели""" + if model_path is None or not os.path.exists(model_path): + print("⚠️ Модель не найдена, создаем новую UNet с случайными весами") + # Создаем модель с случайными весами для тестирования + model = UNet(in_channels=3, out_classes=1) + model.to(self.device) + model.eval() + return model + + try: + model = UNet(in_channels=3, out_classes=1) + state_dict = torch.load(model_path, map_location=self.device) + model.load_state_dict(state_dict) + model.to(self.device) + model.eval() + print(f"✅ Модель загружена из {model_path}") + return model + except Exception as e: + print(f"❌ Ошибка загрузки модели: {e}") + # Создаем заглушку + model = UNet(in_channels=3, out_classes=1) + model.to(self.device) + model.eval() + return model + + def _load_totalsegmentator(self): + """Загрузка TotalSegmentator (заглушка)""" + print("ℹ️ TotalSegmentator пока не реализован, используем UNet") + return self._load_unet() + + def _segment_with_unet(self, image: np.ndarray) -> np.ndarray: + """ + Сегментация с помощью UNet + + Args: + image: изображение [H, W] или [H, W, C] + + Returns: + Бинарная маска [H, W] + """ + # Подготовка изображения + if len(image.shape) == 2: + # Если изображение в градациях серого, конвертируем в RGB + image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB) + elif image.shape[-1] == 4: + # Если есть альфа-канал, убираем его + image = image[:, :, :3] + + # Нормализация и ресайз + original_shape = image.shape[:2] + + # Ресайз до размера, который ожидает модель (обычно 256x256 или 512x512) + target_size = 256 + if image.shape[0] != target_size or image.shape[1] != target_size: + image = cv2.resize(image, (target_size, target_size)) + + # Преобразование в тензор + image_tensor = torch.from_numpy(image).float().permute(2, 0, 1) / 255.0 + image_tensor = image_tensor.unsqueeze(0).to(self.device) + + # Инференс + with torch.no_grad(): + output = self.model(image_tensor) + + # Если вывод модели - многоканальный (классы), берем argmax + if output.shape[1] > 1: + mask = torch.argmax(output, dim=1) + else: + # Если бинарная сегментация, применяем порог + mask = torch.sigmoid(output) + mask = (mask > 0.5).float() + + mask = mask.squeeze().cpu().numpy() + + # Ресайз обратно к оригинальному размеру + if mask.shape != original_shape: + mask = cv2.resize(mask.astype(np.float32), (original_shape[1], original_shape[0])) + mask = (mask > 0.5).astype(np.uint8) + + return mask + + def _segment_with_totalsegmentator(self, image: np.ndarray) -> np.ndarray: + """Сегментация с помощью TotalSegmentator (заглушка)""" + print("ℹ️ Используем заглушку TotalSegmentator") + # Создаем искусственную маску (круг в центре) + h, w = image.shape[:2] + mask = np.zeros((h, w), dtype=np.uint8) + center = (w // 2, h // 2) + radius = min(h, w) // 4 + cv2.circle(mask, center, radius, 1, -1) + return mask + + def preprocess(self, image: np.ndarray) -> np.ndarray: + """Предобработка изображения перед сегментацией""" + # Если изображение в градациях серого + if len(image.shape) == 2: + image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB) + + # Нормализация интенсивности + image = image.astype(np.float32) + if image.max() > 1.0: + image = image / 255.0 + + # CLAHE для улучшения контраста (опционально) + if len(image.shape) == 3 and image.shape[2] == 3: + lab = cv2.cvtColor((image * 255).astype(np.uint8), cv2.COLOR_RGB2LAB) + lab[:, :, 0] = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8)).apply(lab[:, :, 0]) + image = cv2.cvtColor(lab, cv2.COLOR_LAB2RGB) / 255.0 + + return image + + +# Создание тестовой модели для отладки +def create_dummy_unet(): + """Создает UNet с случайными весами для тестирования""" + from src import UNet + model = UNet(in_channels=3, out_classes=1) + return model + + +# Функция для создания маски-заглушки (если нет модели) +def create_dummy_mask(image: np.ndarray) -> np.ndarray: + """Создает искусственную маску для тестирования""" + h, w = image.shape[:2] + mask = np.zeros((h, w), dtype=np.uint8) + + # Создаем эллипс в центре + center = (w // 2, h // 2) + axes = (w // 4, h // 3) + cv2.ellipse(mask, center, axes, 0, 0, 360, 1, -1) + + # Добавляем немного шума + noise = np.random.random((h, w)) < 0.02 + mask[noise] = 1 + + return mask \ No newline at end of file diff --git a/src/model/unet.py b/src/model/unet.py new file mode 100644 index 0000000..36911b5 --- /dev/null +++ b/src/model/unet.py @@ -0,0 +1,122 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F + + +class DoubleConv(nn.Module): + """Два сверточных слоя с GELU активацией""" + + def __init__(self, in_channels, out_channels, dropout=0.0): + super().__init__() + self.double_conv = nn.Sequential( + nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1), + nn.BatchNorm2d(out_channels), + nn.GELU(), + nn.Dropout2d(dropout), # добавляем dropout для регуляризации + nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1), + nn.BatchNorm2d(out_channels), + nn.GELU(), + ) + + def forward(self, x): + return self.double_conv(x) + + +class DownBlock(nn.Module): + """Блок энкодера: DoubleConv + MaxPool""" + + def __init__(self, in_channels, out_channels, dropout=0.0): + super().__init__() + self.double_conv = DoubleConv(in_channels, out_channels, dropout) + self.maxpool = nn.MaxPool2d(2) + + def forward(self, x): + # Сохраняем результат DoubleConv для skip connection + conv_out = self.double_conv(x) + pooled_out = self.maxpool(conv_out) + return conv_out, pooled_out + + +class UpBlock(nn.Module): + """Блок декодера: Upsample + Concatenate + DoubleConv""" + + def __init__(self, in_channels, out_channels, dropout=0.0): + super().__init__() + # Транспонированная свертка для апсемплинга + self.up_conv = nn.ConvTranspose2d( + in_channels, out_channels, kernel_size=2, stride=2 + ) + # После конкатенации будет (out_channels * 2) + self.double_conv = DoubleConv(out_channels * 2, out_channels, dropout) + + def forward(self, x, skip_connection): + x = self.up_conv(x) + # Обрезаем skip_connection до размера x (если размеры не совпадают) + diffY = skip_connection.size()[2] - x.size()[2] + diffX = skip_connection.size()[3] - x.size()[3] + x = F.pad(x, [diffX // 2, diffX - diffX // 2, + diffY // 2, diffY - diffY // 2]) + # Конкатенируем + x = torch.cat([skip_connection, x], dim=1) + return self.double_conv(x) + + +class UNet(nn.Module): + """Полная архитектура U-Net""" + + def __init__(self, in_channels=1, out_classes=2, features=[64, 128, 256, 512], dropout=0.1): + super().__init__() + self.encoder = nn.ModuleList() + self.decoder = nn.ModuleList() + + # Encoder path + prev_channels = in_channels + for f in features: + self.encoder.append(DownBlock(prev_channels, f, dropout)) + prev_channels = f + + # Bottleneck + self.bottleneck = DoubleConv(prev_channels, features[-1] * 2, dropout) + + # Decoder path (идем в обратном порядке) + for f in reversed(features): + self.decoder.append(UpBlock(f * 2, f, dropout)) + + # Final classification layer + self.final_conv = nn.Conv2d(features[0], out_classes, kernel_size=1) + + def forward(self, x): + skip_connections = [] + + # Encoder + for down_block in self.encoder: + skip, x = down_block(x) + skip_connections.append(skip) + + # Bottleneck + x = self.bottleneck(x) + + # Decoder + for up_block in self.decoder: + # Берем последний skip connection + skip = skip_connections.pop() + x = up_block(x, skip) + + # Final layer + x = self.final_conv(x) + return x + + +def test_unet(): + """Тестируем модель на случайном тензоре""" + model = UNet(in_channels=1, out_classes=2) + x = torch.randn(1, 1, 512, 512) # batch_size, channels, height, width + with torch.no_grad(): + y = model(x) + print(f"Input shape: {x.shape}") + print(f"Output shape: {y.shape}") + print(f"Number of parameters: {sum(p.numel() for p in model.parameters()):,}") + + +if __name__ == "__main__": + test_unet() \ No newline at end of file diff --git a/src/quality/__init__.py b/src/quality/__init__.py new file mode 100644 index 0000000..4ada364 --- /dev/null +++ b/src/quality/__init__.py @@ -0,0 +1,23 @@ +""" +Оценка качества медицинских изображений +""" + +from src.quality.quality_scorer_old import QualityScorer +from src.quality.artifact_detector import ArtifactDetector +from src.quality.position_validator import PositionValidator + +# Словарь с типами нарушений +QUALITY_ISSUES = { + "ARTIFACT_MOTION": "Артефакт движения", + "ARTIFACT_NOISE": "Шум на изображении", + "POSITION_SHIFT": "Смещение объекта", + "MARKER_ERROR": "Ошибка разметки", + "LOW_CONTRAST": "Низкий контраст", +} + +__all__ = [ + "QualityScorer", + "ArtifactDetector", + "PositionValidator", + "QUALITY_ISSUES", +] \ No newline at end of file diff --git a/src/quality/artifact_detector.py b/src/quality/artifact_detector.py new file mode 100644 index 0000000..6ff7e92 --- /dev/null +++ b/src/quality/artifact_detector.py @@ -0,0 +1,3 @@ +class ArtifactDetector(): + def __int__(self): + pass \ No newline at end of file diff --git a/src/quality/medical_quality.py b/src/quality/medical_quality.py new file mode 100644 index 0000000..8431cbb --- /dev/null +++ b/src/quality/medical_quality.py @@ -0,0 +1,218 @@ +from typing import Dict, Any + +import numpy as np +from monai.transforms.utils import ndimage + +from src import QualityScorer + + +class MedicalQualityScorer(QualityScorer): + def __init__(self, anatomy_type="spine"): + super().__init__() + self.anatomy_type = anatomy_type # "spine", "femur", "whole_body" + + def evaluate(self, image, segmentation) -> Dict[str, Any]: + # Базовые проверки + base = super().evaluate(image, segmentation) + + # Специфические проверки + if self.anatomy_type == "spine": + medical_check = self.check_spine_segmentation(segmentation) + elif self.anatomy_type == "femur": + medical_check = self.check_femur_segmentation(segmentation) + + return { + **base, + "medical_validation": medical_check + } + + def check_spine_segmentation(self, segmentation: np.ndarray) -> Dict[str, Any]: + """ + Комплексная проверка сегментации позвоночника + + Args: + segmentation: бинарная маска [H, W] + + Returns: + Словарь с результатами проверки + """ + results = { + "valid": True, + "issues": [], + "metrics": {}, + "recommendations": [] + } + + # 1. Проверка количества позвонков + labeled, num_features = ndimage.label(segmentation) + results["metrics"]["num_vertebrae"] = num_features + + if num_features < 3: + results["valid"] = False + results["issues"].append({ + "type": "TOO_FEW_VERTEBRAE", + "details": f"Найдено только {num_features} позвонков, ожидается минимум 3", + "severity": "HIGH" + }) + results["recommendations"].append("Повторите исследование, захватите больше позвонков") + return results # Дальнейшие проверки не имеют смысла + + # 2. Собираем центры позвонков + centers = [] + sizes = [] + for i in range(1, num_features + 1): + y, x = np.where(labeled == i) + if len(y) > 0: + centers.append((np.mean(y), np.mean(x))) + sizes.append(len(y)) + + # 3. Проверка выравнивания + x_coords = [c[1] for c in centers] + x_std = np.std(x_coords) + results["metrics"]["alignment_deviation"] = float(x_std) + + if x_std > 20: + results["valid"] = False + results["issues"].append({ + "type": "POOR_ALIGNMENT", + "details": f"Позвонки не выровнены (отклонение {x_std:.1f} пикселей)", + "severity": "MEDIUM" + }) + results["recommendations"].append("Проверьте укладку пациента") + + # 4. Проверка расстояний между позвонками + if len(centers) >= 2: + centers_sorted = sorted(centers, key=lambda c: c[0]) + distances = [] + for i in range(len(centers_sorted) - 1): + y1, _ = centers_sorted[i] + y2, _ = centers_sorted[i + 1] + distances.append(abs(y2 - y1)) + + if len(distances) > 0: + mean_dist = np.mean(distances) + std_dist = np.std(distances) + results["metrics"]["spacing_variation"] = float(std_dist / mean_dist) + + if std_dist / mean_dist > 0.3: + results["valid"] = False + results["issues"].append({ + "type": "IRREGULAR_SPACING", + "details": f"Неравномерное расстояние между позвонками", + "severity": "MEDIUM" + }) + results["recommendations"].append("Возможно, пропущен позвонок или артефакт") + + # 5. Проверка размера позвонков + if len(sizes) > 1: + mean_size = np.mean(sizes) + results["metrics"]["mean_vertebra_size"] = float(mean_size) + + for i, size in enumerate(sizes): + if abs(size - mean_size) / mean_size > 0.5: + results["valid"] = False + results["issues"].append({ + "type": "ABNORMAL_SIZE", + "details": f"Позвонок {i + 1} аномального размера", + "severity": "HIGH" + }) + results["recommendations"].append("Проверьте сегментацию, возможно артефакт") + break + + # Итоговая оценка + if not results["issues"]: + results["status"] = "GOOD" + results["recommendations"].append("Исследование выполнено качественно") + else: + results["status"] = "WARNING" if len(results["issues"]) <= 2 else "POOR" + + return results + + def check_spine_alignment(self, segmentation): + """Проверка, что позвонки выровнены по вертикали""" + labeled, num_features = ndimage.label(segmentation) + + # Собираем центры всех позвонков + centers = [] + for i in range(1, num_features + 1): + y, x = np.where(labeled == i) + if len(y) > 0: + center_y = np.mean(y) + center_x = np.mean(x) + centers.append((center_y, center_x)) + + # Проверяем, что центры находятся примерно на одной вертикальной линии + x_coords = [c[1] for c in centers] + x_std = np.std(x_coords) # стандартное отклонение по X + + # Если отклонение большое → позвоночник искривлен + # Если позвонки смещены влево-вправо → пациент лежал неправильно. + if x_std > 20: # порог в пикселях + return { + "valid": False, + "issue": "позвонки не выровнены", + "deviation": x_std, + "recommendation": "проверьте укладку пациента" + } + + return {"valid": True, "alignment": "good"} + + def check_vertebrae_spacing(self, centers): + """Проверка равномерности расстояния между позвонками""" + # Сортируем по Y (сверху вниз) + centers_sorted = sorted(centers, key=lambda c: c[0]) + + # Вычисляем расстояния между соседними позвонками + distances = [] + for i in range(len(centers_sorted) - 1): + y1, _ = centers_sorted[i] + y2, _ = centers_sorted[i + 1] + distances.append(abs(y2 - y1)) + + # Если расстояния сильно различаются → проблема + if len(distances) > 0: + mean_dist = np.mean(distances) + std_dist = np.std(distances) + + if std_dist / mean_dist > 0.3: # больше 30% вариации + return { + "valid": False, + "issue": "неравномерное расстояние между позвонками", + "distances": distances, + "variation": std_dist / mean_dist + } + + return {"valid": True} + + def check_vertebrae_size(labeled, num_features): + """Проверка, что все позвонки примерно одного размера""" + sizes = [] + for i in range(1, num_features + 1): + size = np.sum(labeled == i) + sizes.append(size) + + # Если один позвонок сильно отличается по размеру + if len(sizes) > 1: + mean_size = np.mean(sizes) + std_size = np.std(sizes) + + # Проверяем каждый позвонок + issues = [] + for i, size in enumerate(sizes): + if abs(size - mean_size) / mean_size > 0.5: # отклонение > 50% + issues.append({ + "vertebra": i + 1, + "size": size, + "expected": mean_size, + "issue": "позвонок аномального размера" + }) + + if issues: + return { + "valid": False, + "issue": "обнаружены позвонки аномального размера", + "details": issues, + "recommendation": "возможно, артефакт или неправильная сегментация" + } + + return {"valid": True} \ No newline at end of file diff --git a/src/quality/position_validator.py b/src/quality/position_validator.py new file mode 100644 index 0000000..983b6a4 --- /dev/null +++ b/src/quality/position_validator.py @@ -0,0 +1,3 @@ +class PositionValidator(): + def __init__(self): + pass \ No newline at end of file diff --git a/src/quality/quality_scorer.py b/src/quality/quality_scorer.py new file mode 100644 index 0000000..85d2840 --- /dev/null +++ b/src/quality/quality_scorer.py @@ -0,0 +1,260 @@ +import numpy as np +from scipy import ndimage +from typing import Dict, Any, Optional, Tuple + + +def convert_to_serializable(obj): + """Рекурсивно конвертирует numpy типы в Python типы для JSON сериализации""" + if isinstance(obj, np.ndarray): + return obj.tolist() + # Для NumPy 2.0 используем только актуальные типы + elif isinstance(obj, (np.int8, np.int16, np.int32, np.int64, + np.uint8, np.uint16, np.uint32, np.uint64)): + return int(obj) + elif isinstance(obj, (np.float16, np.float32, np.float64)): + return float(obj) + elif isinstance(obj, (np.bool_)): + return bool(obj) + elif isinstance(obj, dict): + return {key: convert_to_serializable(value) for key, value in obj.items()} + elif isinstance(obj, (list, tuple)): + return [convert_to_serializable(item) for item in obj] + else: + return obj + + +class QualityScorer: + """Оценка качества сегментации (работает с любыми объектами)""" + + def __init__(self): + self.quality_metrics = {} + + def check_artifact(self, segmentation: np.ndarray, context: str = "single_object") -> Dict[str, Any]: + """ + Проверка артефактов (движение, шум, размытость) + """ + if segmentation is None or segmentation.size == 0: + return {"artifact": True, "type": "missing", "severity": "critical"} + + if segmentation.sum() == 0: + return {"artifact": True, "type": "empty", "severity": "critical"} + + # 1. Проверка четкости границ + gradient = np.gradient(segmentation.astype(float)) + edge_energy = float(np.sqrt(gradient[0] ** 2 + gradient[1] ** 2).mean()) + + # 2. Проверка компактности объекта + labeled, num_features = ndimage.label(segmentation) + num_features = int(num_features) + + # 3. Проверка размера объекта + object_size = float(segmentation.sum() / segmentation.size) + + # 4. Проверка размытости + hist, _ = np.histogram(segmentation, bins=2) + total_pixels = segmentation.size + if total_pixels > 0: + p = hist / total_pixels + p = p[p > 0] + entropy = float(-np.sum(p * np.log2(p + 1e-10))) + else: + entropy = 0.0 + + # Логика для разных контекстов + if context == "multiple_objects": + if num_features < 3: + return { + "artifact": True, + "type": "too_few_objects", + "severity": "high", + "num_objects": num_features, + "edge_energy": edge_energy, + "object_size": object_size + } + elif num_features > 7: + return { + "artifact": True, + "type": "too_many_objects", + "severity": "medium", + "num_objects": num_features, + "edge_energy": edge_energy, + "object_size": object_size + } + else: + if num_features > 1: + sizes = [np.sum(labeled == i) for i in range(1, num_features + 1)] + main_object_size = max(sizes) + total_size = sum(sizes) + + if main_object_size / total_size > 0.8: + return { + "artifact": False, + "num_objects": num_features, + "minor_objects": num_features - 1, + "edge_energy": edge_energy, + "object_size": object_size + } + else: + return { + "artifact": True, + "type": "fragmented", + "severity": "high", + "num_objects": num_features, + "edge_energy": edge_energy, + "object_size": object_size + } + + return { + "artifact": False, + "num_objects": num_features, + "edge_energy": edge_energy, + "object_size": object_size, + "entropy": entropy + } + + def check_position(self, segmentation: np.ndarray, + expected_center: Tuple[float, float] = (0.5, 0.5), + allowed_deviation: float = 0.25) -> Dict[str, Any]: + """ + Проверка положения объекта в кадре + """ + if segmentation is None or segmentation.sum() == 0: + return {"position": None, "valid": False, "error": "empty_mask"} + + y, x = np.where(segmentation > 0) + if len(y) == 0: + return {"position": None, "valid": False, "error": "no_pixels"} + + center_y, center_x = np.mean(y), np.mean(x) + h, w = segmentation.shape + + norm_y = float(center_y / h) + norm_x = float(center_x / w) + + deviation = float(np.sqrt((norm_y - expected_center[0]) ** 2 + + (norm_x - expected_center[1]) ** 2)) + + object_size = float(len(y) / (h * w)) + valid = bool(deviation < allowed_deviation and object_size > 0.005) + + return { + "position": (norm_y, norm_x), + "valid": valid, + "deviation": deviation, + "object_size": object_size, + "center_pixels": (int(center_y), int(center_x)), + "expected_center": expected_center, + "allowed_deviation": allowed_deviation + } + + def check_contrast(self, image: np.ndarray, segmentation: np.ndarray) -> Dict[str, Any]: + """ + Проверка контрастности области объекта + """ + if segmentation is None or segmentation.sum() == 0: + return {"valid": False, "contrast": 0.0, "error": "empty_mask"} + + # Проверка размеров + if image.shape[:2] != segmentation.shape: + from PIL import Image + seg_pil = Image.fromarray(segmentation.astype(np.uint8)) + seg_pil = seg_pil.resize((image.shape[1], image.shape[0]), resample=Image.NEAREST) + segmentation = np.array(seg_pil) + + # Конвертация в grayscale + if len(image.shape) == 3: + image_gray = np.mean(image, axis=2) + else: + image_gray = image + + object_pixels = image_gray[segmentation > 0] + if len(object_pixels) == 0: + return {"valid": False, "contrast": 0.0, "error": "no_pixels"} + + background_pixels = image_gray[segmentation == 0] + if len(background_pixels) == 0: + background_mean = 0.0 + else: + background_mean = float(np.mean(background_pixels)) + + contrast = float(np.std(object_pixels) / (np.mean(object_pixels) + 1e-8)) + mean_ratio = float(np.mean(object_pixels) / (np.mean(background_pixels) + 1e-8)) + mean_intensity = float(np.mean(object_pixels)) + std_intensity = float(np.std(object_pixels)) + + return { + "valid": bool(contrast > 0.08), + "contrast": contrast, + "mean_ratio": mean_ratio, + "mean_intensity": mean_intensity, + "std_intensity": std_intensity, + "background_mean": background_mean + } + + def evaluate(self, image: Optional[np.ndarray] = None, + segmentation: Optional[np.ndarray] = None) -> Dict[str, Any]: + """ + Полная оценка качества + """ + results = { + "overall_quality": "GOOD", + "severity": "LOW", + "issues": [], + "metrics": {}, + "confidence": 0.85 + } + + # 1. Проверка артефактов + artifact_results = self.check_artifact(segmentation) + results["metrics"]["artifact"] = artifact_results + + if artifact_results.get("artifact", False): + results["issues"].append({ + "type": "ARTIFACT", + "details": artifact_results.get("type", "unknown"), + "severity": artifact_results.get("severity", "MEDIUM") + }) + results["severity"] = max(results["severity"], + artifact_results.get("severity", "MEDIUM")) + + # 2. Проверка позиции + position_results = self.check_position(segmentation) + results["metrics"]["position"] = position_results + + if not position_results.get("valid", False): + results["issues"].append({ + "type": "POSITION_ERROR", + "details": f"Object at {position_results.get('position')}, deviation {position_results.get('deviation', 0):.3f}", + "severity": "HIGH" + }) + results["severity"] = "HIGH" + + # 3. Проверка контраста + if image is not None: + contrast_results = self.check_contrast(image, segmentation) + results["metrics"]["contrast"] = contrast_results + + if not contrast_results.get("valid", False): + results["issues"].append({ + "type": "LOW_CONTRAST", + "details": f"Contrast: {contrast_results.get('contrast', 0):.3f}", + "severity": "MEDIUM" + }) + if results["severity"] != "HIGH": + results["severity"] = "MEDIUM" + + # Интегральная оценка + if not results["issues"]: + results["overall_quality"] = "GOOD" + results["confidence"] = 0.9 + elif results["severity"] == "HIGH": + results["overall_quality"] = "POOR" + results["confidence"] = 0.5 + else: + results["overall_quality"] = "WARNING" + results["confidence"] = 0.7 + + # Конвертируем все numpy типы + results = convert_to_serializable(results) + + return results \ No newline at end of file diff --git a/src/quality/quality_scorer_old.py b/src/quality/quality_scorer_old.py new file mode 100644 index 0000000..07a5ef7 --- /dev/null +++ b/src/quality/quality_scorer_old.py @@ -0,0 +1,249 @@ +import numpy as np +from scipy import ndimage +from typing import Dict, Any, Optional, Tuple + + +class QualityScorer: + """Оценка качества сегментации (работает с любыми объектами)""" + + def __init__(self): + self.quality_metrics = {} + + def check_artifact(self, segmentation: np.ndarray, context: str = "single_object") -> Dict[str, Any]: + """ + Проверка артефактов с учетом контекста + + Args: + segmentation: бинарная маска + context: "single_object" (кот/собака) или "multiple_objects" (позвонки) + """ + if segmentation is None or segmentation.sum() == 0: + return {"artifact": True, "type": "empty", "severity": "critical"} + + # Находим все объекты + labeled, num_features = ndimage.label(segmentation) + + # Разная логика для разных контекстов + if context == "multiple_objects": + # Для позвонков: ожидаем 3-5 объектов + if num_features < 3: + return {"artifact": True, "type": "too_few_objects", "severity": "high"} + elif num_features > 7: + return {"artifact": True, "type": "too_many_objects", "severity": "medium"} + else: + # Для кошек/собак (или бедренной кости): ожидаем 1 объект + if num_features > 1: + # Проверяем размер объектов + sizes = [np.sum(labeled == i) for i in range(1, num_features + 1)] + main_object_size = max(sizes) + total_size = sum(sizes) + + # Если один объект намного больше остальных - это шум + if main_object_size / total_size > 0.8: + return {"artifact": False, "num_objects": num_features, + "minor_objects": num_features - 1} + else: + return {"artifact": True, "type": "fragmented", "severity": "high"} + + return { + "artifact": False, + "num_objects": num_features, + "object_sizes": [np.sum(labeled == i) for i in range(1, num_features + 1)] + } + + def check_confidence(self, segmentation: np.ndarray, probability_map: np.ndarray) -> Dict[str, Any]: + """ + Проверка уверенности модели в сегментации + """ + # Если модель не уверена - это может быть артефакт + mean_confidence = probability_map[segmentation > 0].mean() + + return { + "confidence": float(mean_confidence), + "valid": mean_confidence > 0.5, # порог уверенности + "low_confidence_regions": (probability_map < 0.3).sum() / segmentation.sum() + } + + def check_position(self, segmentation: np.ndarray, + expected_center: Tuple[float, float] = (0.5, 0.5), + allowed_deviation: float = 0.25) -> Dict[str, Any]: + """ + Проверка положения объекта в кадре + + Args: + segmentation: бинарная маска сегментации [H, W] + expected_center: ожидаемый центр (y, x) в нормализованных координатах + allowed_deviation: максимальное допустимое отклонение + + Returns: + Словарь с результатами проверки + """ + if segmentation is None or segmentation.sum() == 0: + return {"position": None, "valid": False, "error": "empty_mask"} + + # Центр масс + y, x = np.where(segmentation > 0) + if len(y) == 0: + return {"position": None, "valid": False, "error": "no_pixels"} + + center_y, center_x = np.mean(y), np.mean(x) + h, w = segmentation.shape + + # Нормализованные координаты (0-1) + norm_y = center_y / h + norm_x = center_x / w + + # Отклонение от центра + deviation = np.sqrt((norm_y - expected_center[0]) ** 2 + + (norm_x - expected_center[1]) ** 2) + + # Объект должен быть в центре + valid = deviation < allowed_deviation + + # Проверка размера (не слишком маленький) + object_size = len(y) / (h * w) + if object_size < 0.005: # меньше 0.5% площади + valid = False + + return { + "position": (float(norm_y), float(norm_x)), + "valid": valid, + "deviation": float(deviation), + "object_size": float(object_size), + "center_pixels": (int(center_y), int(center_x)), + "expected_center": expected_center, + "allowed_deviation": allowed_deviation + } + + def check_contrast(self, image: np.ndarray, segmentation: np.ndarray) -> Dict[str, Any]: + """ + Проверка контрастности области объекта + + Args: + image: исходное изображение [H, W] или [H, W, C] + segmentation: бинарная маска сегментации [H, W] + + Returns: + Словарь с результатами проверки + """ + if segmentation is None or segmentation.sum() == 0: + return {"valid": False, "contrast": 0, "error": "empty_mask"} + + # Если изображение цветное, конвертируем в grayscale + if len(image.shape) == 3: + image_gray = np.mean(image, axis=2) + else: + image_gray = image + + # Значения пикселей внутри объекта + object_pixels = image_gray[segmentation > 0] + if len(object_pixels) == 0: + return {"valid": False, "contrast": 0, "error": "no_pixels"} + + # Значения пикселей снаружи объекта (фон) + background_pixels = image_gray[segmentation == 0] + + # Контраст: отношение std/mean внутри объекта + contrast = np.std(object_pixels) / (np.mean(object_pixels) + 1e-8) + + # Отношение средних (объект/фон) + mean_ratio = np.mean(object_pixels) / (np.mean(background_pixels) + 1e-8) + + return { + "valid": contrast > 0.08, # порог + "contrast": float(contrast), + "mean_ratio": float(mean_ratio), + "mean_intensity": float(np.mean(object_pixels)), + "std_intensity": float(np.std(object_pixels)), + "background_mean": float(np.mean(background_pixels)) + } + + def evaluate(self, image: Optional[np.ndarray] = None, + segmentation: Optional[np.ndarray] = None) -> Dict[str, Any]: + """ + Полная оценка качества + + Args: + image: исходное изображение [H, W] или [H, W, C] (опционально) + segmentation: бинарная маска сегментации [H, W] + + Returns: + Словарь с полной оценкой качества + """ + results = { + "overall_quality": "GOOD", + "severity": "LOW", + "issues": [], + "metrics": {}, + "confidence": 0.85, + "position": {}, + "artifact": {} + } + + # 1. Проверка артефактов + artifact_results = self.check_artifact(segmentation) + results["metrics"]["artifact"] = artifact_results + + if artifact_results.get("artifact", False): + results["issues"].append({ + "type": "ARTIFACT", + "details": artifact_results.get("type", "unknown"), + "severity": artifact_results.get("severity", "MEDIUM") + }) + results["severity"] = max(results["severity"], + artifact_results.get("severity", "MEDIUM")) + + # 2. Проверка позиции + position_results = self.check_position(segmentation) + results["metrics"]["position"] = position_results + + if not position_results.get("valid", False): + results["issues"].append({ + "type": "POSITION_ERROR", + "details": f"Object at {position_results.get('position')}, deviation {position_results.get('deviation', 0):.3f}", + "severity": "HIGH" + }) + results["severity"] = "HIGH" + + # 3. Проверка контраста (если есть изображение) + if image is not None: + contrast_results = self.check_contrast(image, segmentation) + results["metrics"]["contrast"] = contrast_results + + if not contrast_results.get("valid", False): + results["issues"].append({ + "type": "LOW_CONTRAST", + "details": f"Contrast: {contrast_results.get('contrast', 0):.3f}", + "severity": "MEDIUM" + }) + if results["severity"] != "HIGH": + results["severity"] = "MEDIUM" + + # Интегральная оценка + if not results["issues"]: + results["overall_quality"] = "GOOD" + results["confidence"] = 0.9 + elif results["severity"] == "HIGH": + results["overall_quality"] = "POOR" + results["confidence"] = 0.5 + else: + results["overall_quality"] = "WARNING" + results["confidence"] = 0.7 + + return results + + +# # В медицинском сервисе +# segmentation = model.segment(dicom_image) +# spine_check = check_spine_segmentation(segmentation) +# +# if spine_check["valid"]: +# print("✅ Позвоночник сегментирован правильно") +# print(f" Найдено позвонков: {spine_check['metrics']['num_vertebrae']}") +# print(f" Выравнивание: {spine_check['metrics']['alignment_deviation']:.1f} пикселей") +# else: +# print("❌ Обнаружены проблемы с сегментацией позвоночника:") +# for issue in spine_check["issues"]: +# print(f" • {issue['type']}: {issue['details']}") +# for rec in spine_check["recommendations"]: +# print(f" 💡 {rec}") \ No newline at end of file diff --git a/train.py b/train.py new file mode 100644 index 0000000..a3e70ee --- /dev/null +++ b/train.py @@ -0,0 +1,285 @@ +# train.py +import torch +import torch.nn as nn +import torch.optim as optim +import numpy as np +from pathlib import Path +import matplotlib.pyplot as plt +from tqdm import tqdm +import time + +from src import UNet +from src.dataloaders.pet_dataset import PetDataset, create_dataloaders + + +def get_device(): + """Автоматический выбор устройства""" + if torch.backends.mps.is_available(): + device = torch.device("mps") + print(f"✅ Используем MPS (Apple Silicon) - {torch.backends.mps.is_built()}") + elif torch.cuda.is_available(): + device = torch.device("cuda") + print(f"✅ Используем CUDA - {torch.cuda.get_device_name(0)}") + else: + device = torch.device("cpu") + print("⚠️ Используем CPU (медленно)") + + # Проверка производительности + if device.type == "mps": + # MPS иногда тормозит на некоторых операциях, проверяем + test_tensor = torch.randn(1000, 1000).to(device) + start = time.time() + _ = test_tensor @ test_tensor.T + elapsed = time.time() - start + print(f" MPS скорость: {elapsed:.4f} сек для matmul 1000x1000") + + return device + + +class DiceLoss(nn.Module): + """Dice Loss для сегментации""" + + def __init__(self, smooth=1e-6): + super().__init__() + self.smooth = smooth + + def forward(self, pred, target): + pred = torch.softmax(pred, dim=1) + target_one_hot = torch.nn.functional.one_hot(target, num_classes=pred.shape[1]) + target_one_hot = target_one_hot.permute(0, 3, 1, 2).float() + + intersection = (pred * target_one_hot).sum(dim=(2, 3)) + union = pred.sum(dim=(2, 3)) + target_one_hot.sum(dim=(2, 3)) + dice = (2. * intersection + self.smooth) / (union + self.smooth) + + return 1 - dice.mean() + + +class Trainer: + def __init__(self, model, train_loader, val_loader, device, + learning_rate=1e-4, use_amp=False): + self.model = model.to(device) + self.train_loader = train_loader + self.val_loader = val_loader + self.device = device + self.use_amp = use_amp and device.type == "cuda" + + # Потери + self.criterion_ce = nn.CrossEntropyLoss() + self.criterion_dice = DiceLoss() + + # Оптимизатор + self.optimizer = optim.AdamW(model.parameters(), lr=learning_rate, + weight_decay=1e-5) + self.scheduler = optim.lr_scheduler.ReduceLROnPlateau( + self.optimizer, mode='min', factor=0.5, patience=5 + ) + + # AMP + if self.use_amp: + self.scaler = torch.cuda.amp.GradScaler() + + # История + self.train_losses = [] + self.val_losses = [] + self.val_dice_scores = [] + + print(f"✅ Trainer инициализирован на {device}") + + def train_epoch(self): + self.model.train() + total_loss = 0 + + for images, masks in tqdm(self.train_loader, desc='Training'): + images = images.to(self.device) + masks = masks.to(self.device) + + self.optimizer.zero_grad() + + if self.use_amp: + with torch.cuda.amp.autocast(): + outputs = self.model(images) + loss_ce = self.criterion_ce(outputs, masks) + loss_dice = self.criterion_dice(outputs, masks) + loss = loss_ce + loss_dice + + self.scaler.scale(loss).backward() + self.scaler.step(self.optimizer) + self.scaler.update() + else: + outputs = self.model(images) + loss_ce = self.criterion_ce(outputs, masks) + loss_dice = self.criterion_dice(outputs, masks) + loss = loss_ce + loss_dice + + loss.backward() + self.optimizer.step() + + total_loss += loss.item() + + return total_loss / len(self.train_loader) + + def validate(self): + self.model.eval() + total_loss = 0 + dice_scores = [] + + with torch.no_grad(): + for images, masks in tqdm(self.val_loader, desc='Validation'): + images = images.to(self.device) + masks = masks.to(self.device) + + outputs = self.model(images) + + # Потери + loss_ce = self.criterion_ce(outputs, masks) + loss_dice = self.criterion_dice(outputs, masks) + loss = loss_ce + loss_dice + total_loss += loss.item() + + # Исправленный Dice Score + pred = torch.softmax(outputs, dim=1) + pred_mask = pred.argmax(dim=1) + + # Вычисляем Dice правильно + dice = self.compute_dice_score(pred_mask, masks) + dice_scores.append(dice) + + avg_loss = total_loss / len(self.val_loader) + avg_dice = np.mean(dice_scores) + + return avg_loss, avg_dice + + def compute_dice_score(self, pred, target): + """ + Правильный Dice Score + pred: [B, H, W] - предсказанные маски (0 или 1) + target: [B, H, W] - истинные маски (0 или 1) + """ + smooth = 1e-6 + + # Преобразуем в float + pred = pred.float() + target = target.float() + + # Пересечение + intersection = (pred * target).sum(dim=(1, 2)) + + # Суммы + pred_sum = pred.sum(dim=(1, 2)) + target_sum = target.sum(dim=(1, 2)) + + # Dice = 2 * |A∩B| / (|A| + |B|) + dice = (2. * intersection + smooth) / (pred_sum + target_sum + smooth) + + return dice.mean().item() + + def train(self, epochs, save_best=True): + best_dice = 0 + + for epoch in range(epochs): + print(f"\n{'=' * 50}") + print(f"Epoch {epoch + 1}/{epochs}") + print(f"{'=' * 50}") + + # Train + train_loss = self.train_epoch() + self.train_losses.append(train_loss) + + # Validate + val_loss, val_dice = self.validate() + self.val_losses.append(val_loss) + self.val_dice_scores.append(val_dice) + + # Scheduler + self.scheduler.step(val_loss) + + current_lr = self.optimizer.param_groups[0]['lr'] + + print(f"\n📊 Результаты:") + print(f" Train Loss: {train_loss:.4f}") + print(f" Val Loss: {val_loss:.4f}") + print(f" Val Dice: {val_dice:.4f}") # Теперь должно быть 0-1 + print(f" LR: {current_lr:.6f}") + + # Сохраняем лучшую модель + if save_best and val_dice > best_dice: + best_dice = val_dice + torch.save(self.model.state_dict(), 'best_model_old.pth') + print(f" ✅ Saved best model (Dice: {val_dice:.4f})") + + def plot_history(self): + fig, axes = plt.subplots(1, 2, figsize=(12, 4)) + + axes[0].plot(self.train_losses, label='Train Loss') + axes[0].plot(self.val_losses, label='Val Loss') + axes[0].set_xlabel('Epoch') + axes[0].set_ylabel('Loss') + axes[0].set_title('Training History') + axes[0].legend() + axes[0].grid(True) + + axes[1].plot(self.val_dice_scores, label='Val Dice', color='green') + axes[1].set_xlabel('Epoch') + axes[1].set_ylabel('Dice Score') + axes[1].set_title('Validation Dice Score') + axes[1].legend() + axes[1].grid(True) + + plt.tight_layout() + plt.savefig('training_history.png', dpi=150) + plt.show() + +def main(): + # Пути + data_root = "/Users/denis/workspace/bone_2026/datasets/oxford-iiit-pet" + model_save_path = "models/unet_cats_dogs.pth" + + # Создаем папку для моделей + Path("models").mkdir(exist_ok=True) + + # Определяем устройство + device = get_device() + + # Настройки для быстрого обучения на MPS + batch_size = 16 if device.type == "mps" else 8 # MPS может больше + input_size = (256, 256) + epochs = 20 + + print(f"\n📋 Параметры обучения:") + print(f" Batch size: {batch_size}") + print(f" Input size: {input_size}") + print(f" Epochs: {epochs}") + print(f" Device: {device}") + + # Загрузка данных + print("\n📂 Загрузка данных...") + train_loader, val_loader = create_dataloaders( + data_root=data_root, + batch_size=batch_size, + input_size=input_size + ) + + # Модель + model = UNet(in_channels=3, out_classes=2) + + # Тренировка + trainer = Trainer( + model, + train_loader, + val_loader, + device=device, + learning_rate=1e-4 + ) + trainer.train(epochs=epochs) + + # Сохраняем модель + torch.save(model.state_dict(), model_save_path) + print(f"\n✅ Model saved to {model_save_path}") + + # Визуализация + trainer.plot_history() + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/training_history.png b/training_history.png new file mode 100644 index 0000000..9134f99 Binary files /dev/null and b/training_history.png differ