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