develop
This commit is contained in:
commit
a5ada071dc
|
|
@ -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
|
||||
|
|
@ -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"]
|
||||
|
|
@ -0,0 +1,256 @@
|
|||
# 🦴 Bone Quality Assessment
|
||||
|
||||
[](https://www.python.org/)
|
||||
[](https://fastapi.tiangolo.com/)
|
||||
[](https://pytorch.org/)
|
||||
[](https://www.docker.com/)
|
||||
|
||||
## 📋 Описание
|
||||
|
||||
Сервис искусственного интеллекта для автоматизированной оценки качества денситометрических изображений и их разметки. Система получает на вход рентгеновское денситометрическое исследование в формате DICOM, и оценивает качество выполнения исследования по стандартным критериям, а также корректность разметки анатомических структур на изображениях.
|
||||
|
||||
### Основные возможности
|
||||
|
||||
- 🖼️ **Анализ изображений** — загрузка и обработка медицинских изображений
|
||||
- 🧠 **Сегментация объектов** — выделение анатомических структур (позвонки, кости)
|
||||
- 📊 **Оценка качества** — проверка по 3 критериям:
|
||||
- Артефакты (движение, шум, размытость)
|
||||
- Позиционирование (правильное расположение объекта)
|
||||
- Контрастность (качество изображения)
|
||||
- 🔍 **Детекция нарушений** — определение типа нарушения для некачественных исследований
|
||||
- 🌐 **Web-интерфейс** — удобная загрузка и визуализация результатов
|
||||
- 📡 **REST API** — интеграция с внешними системами
|
||||
- 📈 **Визуализация маски** — отображение сегментации на изображении
|
||||
|
||||
## 🏗️ Архитектура
|
||||
|
||||

|
||||
|
||||
|
||||
## 🚀 Быстрый старт
|
||||
|
||||
### Требования
|
||||
|
||||
- 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
|
||||
|
||||
<div align="center"> <sub>Built with ❤️ for the Bone Quality Assessment Hackathon</sub> </div>
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 32 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 494 KiB |
|
|
@ -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()
|
||||
|
|
@ -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/
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 538 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 458 KiB |
|
|
@ -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()
|
||||
|
|
@ -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
|
||||
|
|
@ -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())
|
||||
|
|
@ -0,0 +1,10 @@
|
|||
"""
|
||||
REST API для сервиса оценки качества
|
||||
"""
|
||||
|
||||
from src.api.endpoints import app
|
||||
|
||||
# Можно добавить middleware или настройки
|
||||
__all__ = [
|
||||
"app",
|
||||
]
|
||||
|
|
@ -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"
|
||||
)
|
||||
|
|
@ -0,0 +1,875 @@
|
|||
<!DOCTYPE html>
|
||||
<html lang="ru">
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||
<title>Bone Quality Assessment - Анализ качества костной ткани</title>
|
||||
<style>
|
||||
* {
|
||||
margin: 0;
|
||||
padding: 0;
|
||||
box-sizing: border-box;
|
||||
}
|
||||
|
||||
body {
|
||||
font-family: 'Segoe UI', Tahoma, Geneva, Verdana, sans-serif;
|
||||
background: linear-gradient(135deg, #667eea 0%, #764ba2 100%);
|
||||
min-height: 100vh;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
padding: 20px;
|
||||
}
|
||||
|
||||
.container {
|
||||
background: white;
|
||||
border-radius: 24px;
|
||||
box-shadow: 0 20px 60px rgba(0, 0, 0, 0.3);
|
||||
padding: 40px;
|
||||
max-width: 1000px;
|
||||
width: 100%;
|
||||
}
|
||||
|
||||
h1 {
|
||||
color: #2d3748;
|
||||
font-size: 28px;
|
||||
margin-bottom: 8px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 12px;
|
||||
}
|
||||
|
||||
.subtitle {
|
||||
color: #718096;
|
||||
font-size: 16px;
|
||||
margin-bottom: 30px;
|
||||
}
|
||||
|
||||
.upload-area {
|
||||
border: 3px dashed #e2e8f0;
|
||||
border-radius: 16px;
|
||||
padding: 50px 20px;
|
||||
text-align: center;
|
||||
transition: all 0.3s ease;
|
||||
cursor: pointer;
|
||||
background: #f7fafc;
|
||||
position: relative;
|
||||
}
|
||||
|
||||
.upload-area:hover {
|
||||
border-color: #667eea;
|
||||
background: #edf2f7;
|
||||
}
|
||||
|
||||
.upload-area.drag-over {
|
||||
border-color: #667eea;
|
||||
background: #eef2ff;
|
||||
transform: scale(1.02);
|
||||
}
|
||||
|
||||
.upload-area .icon {
|
||||
font-size: 48px;
|
||||
margin-bottom: 16px;
|
||||
display: block;
|
||||
}
|
||||
|
||||
.upload-area h3 {
|
||||
color: #2d3748;
|
||||
font-size: 20px;
|
||||
margin-bottom: 8px;
|
||||
}
|
||||
|
||||
.upload-area p {
|
||||
color: #718096;
|
||||
font-size: 14px;
|
||||
}
|
||||
|
||||
.upload-area input[type="file"] {
|
||||
position: absolute;
|
||||
top: 0;
|
||||
left: 0;
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
opacity: 0;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
#previewContainer {
|
||||
display: none;
|
||||
margin-top: 24px;
|
||||
padding: 20px;
|
||||
background: #f7fafc;
|
||||
border-radius: 12px;
|
||||
text-align: center;
|
||||
}
|
||||
|
||||
#previewContainer img {
|
||||
max-width: 100%;
|
||||
max-height: 400px;
|
||||
border-radius: 8px;
|
||||
box-shadow: 0 4px 12px rgba(0, 0, 0, 0.1);
|
||||
}
|
||||
|
||||
#fileName {
|
||||
margin-top: 12px;
|
||||
color: #2d3748;
|
||||
font-weight: 500;
|
||||
}
|
||||
|
||||
.btn-analyze {
|
||||
background: linear-gradient(135deg, #667eea 0%, #764ba2 100%);
|
||||
color: white;
|
||||
border: none;
|
||||
padding: 14px 40px;
|
||||
border-radius: 12px;
|
||||
font-size: 18px;
|
||||
font-weight: 600;
|
||||
cursor: pointer;
|
||||
transition: all 0.3s ease;
|
||||
margin-top: 20px;
|
||||
width: 100%;
|
||||
}
|
||||
|
||||
.btn-analyze:hover:not(:disabled) {
|
||||
transform: translateY(-2px);
|
||||
box-shadow: 0 8px 25px rgba(102, 126, 234, 0.4);
|
||||
}
|
||||
|
||||
.btn-analyze:disabled {
|
||||
opacity: 0.6;
|
||||
cursor: not-allowed;
|
||||
}
|
||||
|
||||
.btn-analyze.loading {
|
||||
position: relative;
|
||||
color: transparent;
|
||||
}
|
||||
|
||||
.btn-analyze.loading::after {
|
||||
content: '';
|
||||
position: absolute;
|
||||
top: 50%;
|
||||
left: 50%;
|
||||
width: 24px;
|
||||
height: 24px;
|
||||
margin: -12px 0 0 -12px;
|
||||
border: 3px solid rgba(255, 255, 255, 0.3);
|
||||
border-top-color: white;
|
||||
border-radius: 50%;
|
||||
animation: spin 0.8s linear infinite;
|
||||
}
|
||||
|
||||
@keyframes spin {
|
||||
to { transform: rotate(360deg); }
|
||||
}
|
||||
|
||||
/* Results */
|
||||
#results {
|
||||
display: none;
|
||||
margin-top: 30px;
|
||||
animation: fadeIn 0.5s ease;
|
||||
}
|
||||
|
||||
@keyframes fadeIn {
|
||||
from { opacity: 0; transform: translateY(20px); }
|
||||
to { opacity: 1; transform: translateY(0); }
|
||||
}
|
||||
|
||||
.result-card {
|
||||
background: #f7fafc;
|
||||
border-radius: 16px;
|
||||
padding: 24px;
|
||||
margin-bottom: 16px;
|
||||
}
|
||||
|
||||
.result-card .label {
|
||||
color: #718096;
|
||||
font-size: 14px;
|
||||
font-weight: 600;
|
||||
text-transform: uppercase;
|
||||
letter-spacing: 0.5px;
|
||||
margin-bottom: 6px;
|
||||
}
|
||||
|
||||
.result-card .value {
|
||||
font-size: 20px;
|
||||
font-weight: 700;
|
||||
color: #2d3748;
|
||||
}
|
||||
|
||||
.quality-badge {
|
||||
display: inline-block;
|
||||
padding: 6px 20px;
|
||||
border-radius: 20px;
|
||||
font-weight: 700;
|
||||
font-size: 18px;
|
||||
}
|
||||
|
||||
.quality-GOOD {
|
||||
background: #c6f6d5;
|
||||
color: #22543d;
|
||||
}
|
||||
|
||||
.quality-WARNING {
|
||||
background: #fefcbf;
|
||||
color: #744210;
|
||||
}
|
||||
|
||||
.quality-POOR {
|
||||
background: #fed7d7;
|
||||
color: #9b2c2c;
|
||||
}
|
||||
|
||||
.severity-LOW {
|
||||
background: #c6f6d5;
|
||||
color: #22543d;
|
||||
}
|
||||
|
||||
.severity-MEDIUM {
|
||||
background: #fefcbf;
|
||||
color: #744210;
|
||||
}
|
||||
|
||||
.severity-HIGH {
|
||||
background: #fed7d7;
|
||||
color: #9b2c2c;
|
||||
}
|
||||
|
||||
.issue-item {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 10px;
|
||||
padding: 8px 0;
|
||||
border-bottom: 1px solid #e2e8f0;
|
||||
}
|
||||
|
||||
.issue-item:last-child {
|
||||
border-bottom: none;
|
||||
}
|
||||
|
||||
.issue-badge {
|
||||
padding: 2px 12px;
|
||||
border-radius: 12px;
|
||||
font-size: 12px;
|
||||
font-weight: 600;
|
||||
}
|
||||
|
||||
.issue-ARTIFACT {
|
||||
background: #fefcbf;
|
||||
color: #744210;
|
||||
}
|
||||
|
||||
.issue-POSITION_ERROR {
|
||||
background: #fed7d7;
|
||||
color: #9b2c2c;
|
||||
}
|
||||
|
||||
.issue-LOW_CONTRAST {
|
||||
background: #bee3f8;
|
||||
color: #2a69ac;
|
||||
}
|
||||
|
||||
.grid-2 {
|
||||
display: grid;
|
||||
grid-template-columns: 1fr 1fr;
|
||||
gap: 16px;
|
||||
margin-top: 16px;
|
||||
}
|
||||
|
||||
.validation-item {
|
||||
background: white;
|
||||
padding: 16px;
|
||||
border-radius: 12px;
|
||||
border-left: 4px solid #667eea;
|
||||
}
|
||||
|
||||
.validation-item .label {
|
||||
font-size: 13px;
|
||||
color: #718096;
|
||||
font-weight: 600;
|
||||
}
|
||||
|
||||
.validation-item .value {
|
||||
font-size: 14px;
|
||||
color: #2d3748;
|
||||
margin-top: 4px;
|
||||
word-break: break-all;
|
||||
}
|
||||
|
||||
.validation-item .detail {
|
||||
font-size: 12px;
|
||||
color: #718096;
|
||||
margin-top: 4px;
|
||||
}
|
||||
|
||||
.confidence-bar {
|
||||
width: 100%;
|
||||
height: 8px;
|
||||
background: #e2e8f0;
|
||||
border-radius: 4px;
|
||||
overflow: hidden;
|
||||
margin-top: 8px;
|
||||
}
|
||||
|
||||
.confidence-bar .fill {
|
||||
height: 100%;
|
||||
background: linear-gradient(90deg, #667eea, #764ba2);
|
||||
border-radius: 4px;
|
||||
transition: width 0.8s ease;
|
||||
}
|
||||
|
||||
.error-message {
|
||||
background: #fed7d7;
|
||||
color: #9b2c2c;
|
||||
padding: 16px;
|
||||
border-radius: 12px;
|
||||
margin-top: 16px;
|
||||
display: none;
|
||||
}
|
||||
|
||||
.file-info {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
gap: 12px;
|
||||
margin-top: 12px;
|
||||
font-size: 14px;
|
||||
color: #4a5568;
|
||||
}
|
||||
|
||||
/* Mask visualization */
|
||||
#maskContainer {
|
||||
display: none;
|
||||
margin-top: 24px;
|
||||
padding: 20px;
|
||||
background: #f7fafc;
|
||||
border-radius: 12px;
|
||||
text-align: center;
|
||||
}
|
||||
|
||||
#maskContainer h3 {
|
||||
color: #2d3748;
|
||||
margin-bottom: 12px;
|
||||
}
|
||||
|
||||
.mask-wrapper {
|
||||
position: relative;
|
||||
display: inline-block;
|
||||
width: 100%;
|
||||
max-width: 100%;
|
||||
overflow: hidden;
|
||||
border-radius: 8px;
|
||||
background: #1a202c;
|
||||
}
|
||||
|
||||
#maskImage {
|
||||
display: block;
|
||||
width: 100%;
|
||||
height: auto;
|
||||
border-radius: 8px;
|
||||
}
|
||||
|
||||
#maskCanvas {
|
||||
position: absolute;
|
||||
top: 0;
|
||||
left: 0;
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
pointer-events: none;
|
||||
border-radius: 8px;
|
||||
}
|
||||
|
||||
.mask-controls {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
gap: 16px;
|
||||
margin: 16px 0;
|
||||
flex-wrap: wrap;
|
||||
}
|
||||
|
||||
.mask-controls label {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
font-size: 14px;
|
||||
color: #4a5568;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.mask-controls input[type="range"] {
|
||||
width: 150px;
|
||||
cursor: pointer;
|
||||
}
|
||||
|
||||
.btn-toggle {
|
||||
background: #edf2f7;
|
||||
border: 1px solid #e2e8f0;
|
||||
padding: 8px 16px;
|
||||
border-radius: 8px;
|
||||
cursor: pointer;
|
||||
font-size: 14px;
|
||||
color: #4a5568;
|
||||
transition: all 0.3s ease;
|
||||
}
|
||||
|
||||
.btn-toggle:hover {
|
||||
background: #e2e8f0;
|
||||
border-color: #667eea;
|
||||
}
|
||||
|
||||
.btn-toggle.active {
|
||||
background: #667eea;
|
||||
color: white;
|
||||
border-color: #667eea;
|
||||
}
|
||||
|
||||
.mask-stats {
|
||||
display: flex;
|
||||
gap: 24px;
|
||||
justify-content: center;
|
||||
flex-wrap: wrap;
|
||||
font-size: 13px;
|
||||
color: #718096;
|
||||
margin-top: 12px;
|
||||
}
|
||||
|
||||
.mask-stats span {
|
||||
background: white;
|
||||
padding: 4px 12px;
|
||||
border-radius: 4px;
|
||||
}
|
||||
|
||||
@media (max-width: 640px) {
|
||||
.container {
|
||||
padding: 20px;
|
||||
}
|
||||
.grid-2 {
|
||||
grid-template-columns: 1fr;
|
||||
}
|
||||
h1 {
|
||||
font-size: 22px;
|
||||
}
|
||||
.mask-controls {
|
||||
flex-direction: column;
|
||||
}
|
||||
.mask-controls input[type="range"] {
|
||||
width: 100%;
|
||||
}
|
||||
}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div class="container">
|
||||
<h1>🦴 Bone Quality Assessment</h1>
|
||||
<p class="subtitle">Загрузите медицинское изображение для анализа качества костной ткани</p>
|
||||
|
||||
<!-- Upload Area -->
|
||||
<div class="upload-area" id="uploadArea">
|
||||
<span class="icon">📤</span>
|
||||
<h3>Перетащите изображение сюда</h3>
|
||||
<p>или кликните для выбора файла</p>
|
||||
<p style="font-size: 12px; color: #a0aec0; margin-top: 8px;">Поддерживаются: PNG, JPG, JPEG</p>
|
||||
<input type="file" id="fileInput" accept=".png,.jpg,.jpeg">
|
||||
</div>
|
||||
|
||||
<!-- Preview -->
|
||||
<div id="previewContainer">
|
||||
<img id="previewImage" src="#" alt="Preview">
|
||||
<div id="fileName"></div>
|
||||
<div class="file-info">
|
||||
<span id="fileSize"></span>
|
||||
<span>•</span>
|
||||
<span id="fileType"></span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<button class="btn-analyze" id="analyzeBtn" disabled>🔍 Анализировать</button>
|
||||
|
||||
<!-- Error -->
|
||||
<div class="error-message" id="errorMessage"></div>
|
||||
|
||||
<!-- Mask Visualization -->
|
||||
<div id="maskContainer">
|
||||
<h3>🧩 Визуализация сегментации</h3>
|
||||
<div class="mask-controls">
|
||||
<label>
|
||||
<span>Прозрачность:</span>
|
||||
<input type="range" id="maskOpacity" min="0" max="100" value="50">
|
||||
<span id="opacityLabel">50%</span>
|
||||
</label>
|
||||
<button class="btn-toggle active" id="toggleMaskBtn">Скрыть маску</button>
|
||||
<button class="btn-toggle" id="toggleOriginalBtn">Показать оригинал</button>
|
||||
</div>
|
||||
<div class="mask-wrapper">
|
||||
<img id="maskImage" src="#" alt="Mask Visualization">
|
||||
<canvas id="maskCanvas"></canvas>
|
||||
</div>
|
||||
<div class="mask-stats" id="maskStats">
|
||||
<span>🧊 Объектов: <strong id="numObjects">-</strong></span>
|
||||
<span>📏 Размер: <strong id="objectSize">-</strong></span>
|
||||
<span>🎯 Уверенность: <strong id="maskConfidence">-</strong></span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Results -->
|
||||
<div id="results">
|
||||
<div class="result-card">
|
||||
<div class="label">Общее качество</div>
|
||||
<div id="overallQuality"></div>
|
||||
</div>
|
||||
|
||||
<div class="result-card">
|
||||
<div class="label">Серьезность проблемы</div>
|
||||
<div id="severityDisplay"></div>
|
||||
</div>
|
||||
|
||||
<div class="result-card">
|
||||
<div class="label">Выявленные проблемы</div>
|
||||
<div id="issuesList"></div>
|
||||
</div>
|
||||
|
||||
<div class="grid-2">
|
||||
<div class="validation-item">
|
||||
<div class="label">📍 Позиционирование</div>
|
||||
<div class="value" id="positionValidation"></div>
|
||||
<div class="detail" id="positionDetails"></div>
|
||||
</div>
|
||||
<div class="validation-item">
|
||||
<div class="label">🔬 Артефакты</div>
|
||||
<div class="value" id="artifactValidation"></div>
|
||||
<div class="detail" id="artifactDetails"></div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="result-card">
|
||||
<div class="label">Уверенность модели</div>
|
||||
<div id="confidenceDisplay"></div>
|
||||
<div class="confidence-bar">
|
||||
<div class="fill" id="confidenceFill" style="width: 0%"></div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="result-card" id="metricsCard" style="display: none;">
|
||||
<div class="label">Детальные метрики</div>
|
||||
<div id="metricsDisplay" style="font-size: 14px; color: #4a5568; max-height: 300px; overflow-y: auto;"></div>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<script>
|
||||
// Конфигурация
|
||||
const API_URL = 'http://localhost:8000/api/v1/analyze';
|
||||
|
||||
// DOM элементы
|
||||
const uploadArea = document.getElementById('uploadArea');
|
||||
const fileInput = document.getElementById('fileInput');
|
||||
const previewContainer = document.getElementById('previewContainer');
|
||||
const previewImage = document.getElementById('previewImage');
|
||||
const fileName = document.getElementById('fileName');
|
||||
const fileSize = document.getElementById('fileSize');
|
||||
const fileType = document.getElementById('fileType');
|
||||
const analyzeBtn = document.getElementById('analyzeBtn');
|
||||
const results = document.getElementById('results');
|
||||
const errorMessage = document.getElementById('errorMessage');
|
||||
const maskContainer = document.getElementById('maskContainer');
|
||||
const maskImage = document.getElementById('maskImage');
|
||||
const maskCanvas = document.getElementById('maskCanvas');
|
||||
const maskOpacity = document.getElementById('maskOpacity');
|
||||
const opacityLabel = document.getElementById('opacityLabel');
|
||||
const toggleMaskBtn = document.getElementById('toggleMaskBtn');
|
||||
const toggleOriginalBtn = document.getElementById('toggleOriginalBtn');
|
||||
const numObjects = document.getElementById('numObjects');
|
||||
const objectSize = document.getElementById('objectSize');
|
||||
const maskConfidence = document.getElementById('maskConfidence');
|
||||
|
||||
// Состояние
|
||||
let selectedFile = null;
|
||||
let currentImageData = null;
|
||||
let currentMaskData = null;
|
||||
let showMask = true;
|
||||
let currentMaskBase64 = null;
|
||||
|
||||
// Drag and drop
|
||||
uploadArea.addEventListener('dragover', (e) => {
|
||||
e.preventDefault();
|
||||
uploadArea.classList.add('drag-over');
|
||||
});
|
||||
|
||||
uploadArea.addEventListener('dragleave', () => {
|
||||
uploadArea.classList.remove('drag-over');
|
||||
});
|
||||
|
||||
uploadArea.addEventListener('drop', (e) => {
|
||||
e.preventDefault();
|
||||
uploadArea.classList.remove('drag-over');
|
||||
if (e.dataTransfer.files.length) {
|
||||
handleFile(e.dataTransfer.files[0]);
|
||||
}
|
||||
});
|
||||
|
||||
fileInput.addEventListener('change', (e) => {
|
||||
if (e.target.files.length) {
|
||||
handleFile(e.target.files[0]);
|
||||
}
|
||||
});
|
||||
|
||||
function handleFile(file) {
|
||||
selectedFile = file;
|
||||
const reader = new FileReader();
|
||||
reader.onload = (e) => {
|
||||
currentImageData = e.target.result;
|
||||
previewImage.src = e.target.result;
|
||||
previewContainer.style.display = 'block';
|
||||
fileName.textContent = file.name;
|
||||
fileSize.textContent = formatFileSize(file.size);
|
||||
fileType.textContent = file.type || 'Неизвестно';
|
||||
analyzeBtn.disabled = false;
|
||||
results.style.display = 'none';
|
||||
maskContainer.style.display = 'none';
|
||||
errorMessage.style.display = 'none';
|
||||
currentMaskData = null;
|
||||
currentMaskBase64 = null;
|
||||
};
|
||||
reader.readAsDataURL(file);
|
||||
}
|
||||
|
||||
function formatFileSize(bytes) {
|
||||
if (bytes < 1024) return bytes + ' B';
|
||||
if (bytes < 1024 * 1024) return (bytes / 1024).toFixed(1) + ' KB';
|
||||
return (bytes / (1024 * 1024)).toFixed(2) + ' MB';
|
||||
}
|
||||
|
||||
// Mask controls
|
||||
maskOpacity.addEventListener('input', () => {
|
||||
const val = maskOpacity.value;
|
||||
opacityLabel.textContent = val + '%';
|
||||
if (currentImageData && currentMaskBase64) {
|
||||
drawMask(currentImageData, currentMaskBase64, val / 100);
|
||||
}
|
||||
});
|
||||
|
||||
toggleMaskBtn.addEventListener('click', () => {
|
||||
showMask = !showMask;
|
||||
toggleMaskBtn.textContent = showMask ? 'Скрыть маску' : 'Показать маску';
|
||||
toggleMaskBtn.classList.toggle('active');
|
||||
if (currentImageData && currentMaskBase64) {
|
||||
drawMask(currentImageData, currentMaskBase64, maskOpacity.value / 100);
|
||||
}
|
||||
});
|
||||
|
||||
toggleOriginalBtn.addEventListener('click', () => {
|
||||
if (currentImageData) {
|
||||
const ctx = maskCanvas.getContext('2d');
|
||||
const img = new Image();
|
||||
img.onload = () => {
|
||||
maskCanvas.width = img.width;
|
||||
maskCanvas.height = img.height;
|
||||
ctx.drawImage(img, 0, 0);
|
||||
};
|
||||
img.src = currentImageData;
|
||||
}
|
||||
toggleMaskBtn.textContent = 'Показать маску';
|
||||
toggleMaskBtn.classList.remove('active');
|
||||
showMask = false;
|
||||
});
|
||||
|
||||
// Draw mask function
|
||||
function drawMask(imageSrc, maskBase64, opacity) {
|
||||
const ctx = maskCanvas.getContext('2d');
|
||||
const img = new Image();
|
||||
img.onload = () => {
|
||||
maskCanvas.width = img.width;
|
||||
maskCanvas.height = img.height;
|
||||
ctx.clearRect(0, 0, maskCanvas.width, maskCanvas.height);
|
||||
ctx.drawImage(img, 0, 0);
|
||||
|
||||
if (showMask && maskBase64) {
|
||||
const maskImg = new Image();
|
||||
maskImg.onload = () => {
|
||||
ctx.globalAlpha = opacity * 0.6;
|
||||
ctx.drawImage(maskImg, 0, 0, maskCanvas.width, maskCanvas.height);
|
||||
ctx.globalAlpha = 1;
|
||||
|
||||
// Информация о маске
|
||||
ctx.fillStyle = 'rgba(255, 255, 255, 0.85)';
|
||||
ctx.fillRect(10, 10, 200, 60);
|
||||
ctx.fillStyle = '#2d3748';
|
||||
ctx.font = '12px Arial';
|
||||
ctx.fillText(`Объектов: ${currentMaskData?.num_objects || 0}`, 20, 30);
|
||||
ctx.fillText(`Размер: ${((currentMaskData?.object_size || 0) * 100).toFixed(1)}%`, 20, 50);
|
||||
ctx.globalAlpha = 1;
|
||||
};
|
||||
maskImg.src = `data:image/png;base64,${maskBase64}`;
|
||||
}
|
||||
};
|
||||
img.src = imageSrc;
|
||||
}
|
||||
|
||||
// Analyze
|
||||
analyzeBtn.addEventListener('click', async () => {
|
||||
if (!selectedFile) return;
|
||||
|
||||
analyzeBtn.disabled = true;
|
||||
analyzeBtn.classList.add('loading');
|
||||
results.style.display = 'none';
|
||||
maskContainer.style.display = 'none';
|
||||
errorMessage.style.display = 'none';
|
||||
|
||||
const formData = new FormData();
|
||||
formData.append('file', selectedFile);
|
||||
|
||||
try {
|
||||
const response = await fetch(API_URL, {
|
||||
method: 'POST',
|
||||
body: formData
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorText = await response.text();
|
||||
throw new Error(`Ошибка сервера (${response.status}): ${errorText}`);
|
||||
}
|
||||
|
||||
const data = await response.json();
|
||||
|
||||
// Сохраняем маску
|
||||
if (data.mask) {
|
||||
currentMaskBase64 = data.mask;
|
||||
currentMaskData = {
|
||||
num_objects: data.metrics?.artifact?.num_objects || 0,
|
||||
object_size: data.metrics?.artifact?.object_size || 0,
|
||||
edge_energy: data.metrics?.artifact?.edge_energy || 0
|
||||
};
|
||||
|
||||
maskContainer.style.display = 'block';
|
||||
maskImage.src = currentImageData;
|
||||
|
||||
// Обновляем статистику
|
||||
numObjects.textContent = currentMaskData.num_objects;
|
||||
objectSize.textContent = (currentMaskData.object_size * 100).toFixed(1) + '%';
|
||||
maskConfidence.textContent = (data.confidence * 100).toFixed(1) + '%';
|
||||
|
||||
// Рисуем маску
|
||||
drawMask(currentImageData, data.mask, maskOpacity.value / 100);
|
||||
}
|
||||
|
||||
displayResults(data);
|
||||
|
||||
} catch (error) {
|
||||
errorMessage.textContent = `❌ Ошибка: ${error.message}`;
|
||||
errorMessage.style.display = 'block';
|
||||
} finally {
|
||||
analyzeBtn.disabled = false;
|
||||
analyzeBtn.classList.remove('loading');
|
||||
}
|
||||
});
|
||||
|
||||
function displayResults(data) {
|
||||
// Overall quality
|
||||
const quality = data.overall_quality || 'GOOD';
|
||||
const qualityEl = document.getElementById('overallQuality');
|
||||
const badgeClass = `quality-${quality}`;
|
||||
const labelMap = {
|
||||
'GOOD': '✅ Хорошее',
|
||||
'WARNING': '⚠️ Требует внимания',
|
||||
'POOR': '❌ Плохое'
|
||||
};
|
||||
qualityEl.innerHTML = `<span class="quality-badge ${badgeClass}">${labelMap[quality] || quality}</span>`;
|
||||
|
||||
// Severity
|
||||
const severity = data.severity || 'LOW';
|
||||
const severityEl = document.getElementById('severityDisplay');
|
||||
const severityClass = `severity-${severity}`;
|
||||
const severityMap = {
|
||||
'LOW': '🟢 Низкая',
|
||||
'MEDIUM': '🟡 Средняя',
|
||||
'HIGH': '🔴 Высокая'
|
||||
};
|
||||
severityEl.innerHTML = `<span class="quality-badge ${severityClass}">${severityMap[severity] || severity}</span>`;
|
||||
|
||||
// Issues
|
||||
const issuesEl = document.getElementById('issuesList');
|
||||
if (data.issues && data.issues.length > 0) {
|
||||
issuesEl.innerHTML = data.issues.map(issue => {
|
||||
const type = issue.type || 'info';
|
||||
const badgeClass = `issue-${type}`;
|
||||
const details = issue.details || '';
|
||||
const severity = issue.severity || '';
|
||||
return `<div class="issue-item">
|
||||
<span class="issue-badge ${badgeClass}">${type.replace('_', ' ')}</span>
|
||||
<span>${details}</span>
|
||||
${severity ? `<span style="font-size: 12px; color: #718096;">(${severity})</span>` : ''}
|
||||
</div>`;
|
||||
}).join('');
|
||||
} else {
|
||||
issuesEl.innerHTML = '<div style="color: #48bb78;">✅ Проблем не обнаружено</div>';
|
||||
}
|
||||
|
||||
// Position validation
|
||||
const pos = data.position_validation || {};
|
||||
const posEl = document.getElementById('positionValidation');
|
||||
const posDetails = document.getElementById('positionDetails');
|
||||
|
||||
if (pos.valid !== undefined) {
|
||||
posEl.innerHTML = pos.valid ? '✅ Корректное' : '❌ Некорректное';
|
||||
posDetails.textContent = pos.position ? `Позиция: ${pos.position}, отклонение: ${(pos.deviation || 0).toFixed(3)}` : '';
|
||||
} else {
|
||||
posEl.textContent = 'Не определено';
|
||||
posDetails.textContent = '';
|
||||
}
|
||||
|
||||
// Artifact validation
|
||||
const art = data.artifact_validation || {};
|
||||
const artEl = document.getElementById('artifactValidation');
|
||||
const artDetails = document.getElementById('artifactDetails');
|
||||
|
||||
if (art.artifact !== undefined) {
|
||||
artEl.innerHTML = art.artifact ? '⚠️ Обнаружены' : '✅ Не обнаружены';
|
||||
artDetails.textContent = art.type ? `Тип: ${art.type}` : '';
|
||||
} else {
|
||||
artEl.textContent = 'Не определено';
|
||||
artDetails.textContent = '';
|
||||
}
|
||||
|
||||
// Confidence
|
||||
const confidence = data.confidence || 0;
|
||||
document.getElementById('confidenceDisplay').textContent =
|
||||
`${(confidence * 100).toFixed(1)}%`;
|
||||
document.getElementById('confidenceFill').style.width = `${(confidence * 100).toFixed(1)}%`;
|
||||
|
||||
// Metrics
|
||||
const metrics = data.metrics || {};
|
||||
const metricsCard = document.getElementById('metricsCard');
|
||||
const metricsDisplay = document.getElementById('metricsDisplay');
|
||||
|
||||
if (Object.keys(metrics).length > 0) {
|
||||
metricsCard.style.display = 'block';
|
||||
metricsDisplay.innerHTML = `<pre style="white-space: pre-wrap; word-break: break-all; font-size: 12px;">${JSON.stringify(metrics, null, 2)}</pre>`;
|
||||
} else {
|
||||
metricsCard.style.display = 'none';
|
||||
}
|
||||
|
||||
results.style.display = 'block';
|
||||
results.scrollIntoView({ behavior: 'smooth', block: 'start' });
|
||||
}
|
||||
|
||||
// Health check
|
||||
async function checkHealth() {
|
||||
try {
|
||||
const response = await fetch('http://localhost:8000/api/v1/health');
|
||||
if (response.ok) {
|
||||
console.log('✅ API сервер работает');
|
||||
} else {
|
||||
console.warn('⚠️ API сервер недоступен');
|
||||
}
|
||||
} catch {
|
||||
console.warn('⚠️ Не удалось подключиться к API серверу');
|
||||
}
|
||||
}
|
||||
checkHealth();
|
||||
|
||||
console.log('🦴 Bone Quality Assessment UI загружен');
|
||||
console.log(`API URL: ${API_URL}`);
|
||||
</script>
|
||||
</body>
|
||||
</html>
|
||||
|
|
@ -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 на хакатоне
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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}")
|
||||
|
|
@ -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
|
||||
|
|
@ -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()
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -0,0 +1,3 @@
|
|||
class ArtifactDetector():
|
||||
def __int__(self):
|
||||
pass
|
||||
|
|
@ -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}
|
||||
|
|
@ -0,0 +1,3 @@
|
|||
class PositionValidator():
|
||||
def __init__(self):
|
||||
pass
|
||||
|
|
@ -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
|
||||
|
|
@ -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}")
|
||||
|
|
@ -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()
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 76 KiB |
Loading…
Reference in New Issue