This commit is contained in:
denis 2026-08-15 23:20:59 +03:00
commit a5ada071dc
30 changed files with 3315 additions and 0 deletions

30
.dockerignore Normal file
View File

@ -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

36
Dockerfile Normal file
View File

@ -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"]

256
README.md Normal file
View File

@ -0,0 +1,256 @@
# 🦴 Bone Quality Assessment
[![Python](https://img.shields.io/badge/Python-3.10-blue.svg)](https://www.python.org/)
[![FastAPI](https://img.shields.io/badge/FastAPI-0.104-green.svg)](https://fastapi.tiangolo.com/)
[![PyTorch](https://img.shields.io/badge/PyTorch-2.1-red.svg)](https://pytorch.org/)
[![Docker](https://img.shields.io/badge/Docker-Ready-blue.svg)](https://www.docker.com/)
## 📋 Описание
Сервис искусственного интеллекта для автоматизированной оценки качества денситометрических изображений и их разметки. Система получает на вход рентгеновское денситометрическое исследование в формате DICOM, и оценивает качество выполнения исследования по стандартным критериям, а также корректность разметки анатомических структур на изображениях.
### Основные возможности
- 🖼️ **Анализ изображений** — загрузка и обработка медицинских изображений
- 🧠 **Сегментация объектов** — выделение анатомических структур (позвонки, кости)
- 📊 **Оценка качества** — проверка по 3 критериям:
- Артефакты (движение, шум, размытость)
- Позиционирование (правильное расположение объекта)
- Контрастность (качество изображения)
- 🔍 **Детекция нарушений** — определение типа нарушения для некачественных исследований
- 🌐 **Web-интерфейс** — удобная загрузка и визуализация результатов
- 📡 **REST API** — интеграция с внешними системами
- 📈 **Визуализация маски** — отображение сегментации на изображении
## 🏗️ Архитектура
![!img](public/static/arch.png)
## 🚀 Быстрый старт
### Требования
- Python 3.10+
- PyTorch 2.1+
- Docker (опционально)
### Локальная установка
```bash
# 1. Клонирование репозитория
git clone https://github.com/yourusername/bone-quality-assessment.git
cd bone-quality-assessment
# 2. Создание виртуального окружения
python -m venv venv
source venv/bin/activate # Linux/Mac
# или
venv\Scripts\activate # Windows
# 3. Установка зависимостей
pip install -r requirements.txt
# 4. Загрузка обученной модели (опционально)
# Поместите модель в папку models/unet_cats_dogs.pth
# 5. Запуск сервера
python run.py
```
Docker
```bash
# 1. Сборка образа
docker build -t bone-quality-api .
# 2. Запуск контейнера
docker run -p 8000:8000 bone-quality-api
# 3. Или используя docker-compose
docker-compose up -d
```
📡 API Endpoints
| Метод | Эндпоинт | Описание |
|--------|-------------------|--------------------------|
| GET | / | Главная страница
| GET | /docs | Swagger UI документация
| GET | /redoc | ReDoc документация
| GET | /api/v1/health | Проверка статуса сервиса
| POST | /api/v1/analyze | Анализ изображения
Пример запроса
```bash
curl -X POST "http://localhost:8000/api/v1/analyze" \
-H "accept: application/json" \
-H "Content-Type: multipart/form-data" \
-F "file=@/path/to/image.jpg"
```
Пример ответа
```json
{
"overall_quality": "GOOD",
"severity": "LOW",
"issues": [],
"confidence": 0.9,
"metrics": {
"artifact": {
"artifact": false,
"num_objects": 1,
"edge_energy": 0.234,
"object_size": 0.123
},
"position": {
"position": [0.45, 0.52],
"valid": true,
"deviation": 0.032
},
"contrast": {
"valid": true,
"contrast": 0.456
}
},
"mask": "base64_encoded_mask_image"
}
```
📊 Интерфейс
Веб-интерфейс доступен по адресу http://localhost:8000/:
- 📤 Drag-and-drop загрузка изображений
- 🔍 Автоматический анализ
- 🎯 Визуализация маски сегментации
- 📈 Детальные метрики качества
- 🏷️ Подробный отчет о нарушениях
🛠️ Технологии
|Компонент | Технология
|-|-|
|Бэкенд | Python 3.10, FastAPI, Uvicorn
|ML | PyTorch, NumPy, SciPy
|Обработка изображений | PIL, OpenCV
|Визуализация | HTML5, CSS3, Canvas API
|Контейнеризация | Docker, Docker Compose
|Документация | Swagger UI, ReDoc
📁 Структура проекта
```text
bone-quality-assessment/
├── src/
│ ├── api/
│ │ ├── endpoints.py # FastAPI эндпоинты
│ │ └── static/
│ │ └── index.html # Web-интерфейс
│ ├── models/
│ │ └── unet.py # U-Net архитектура
│ └── quality/
│ └── quality_scorer.py # Оценка качества
├── models/
│ └── unet_cats_dogs.pth # Обученная модель
├── Dockerfile
├── docker-compose.yml
├── requirements.txt
├── run.py
└── README.md
```
🧪 Тестирование
```bash
# Запуск тестов (если есть)
pytest tests/
```
# Проверка API
```bash
curl http://localhost:8000/api/v1/health
```
## 📊 Метрики качества
Артефакты
*Описание:* Обнаружение шума, размытости, фрагментации
Проверка:
- Энергия границ (edge_energy)
- Количество объектов (num_objects)
- Размер объекта (object_size)
### Позиционирование
*Описание:* Проверка правильности расположения объекта
Проверка:
- Центр масс объекта
- Отклонение от центра изображения
- Минимальный размер объекта
### Контраст
*Описание:* Оценка качества изображения
Проверка:
- Контраст внутри объекта
- Отношение средних (объект/фон)
- Интенсивность пикселей
## 🚀 Деплой
На сервер
```bash
# Копирование на сервер
scp -r ./bone-quality-assessment user@server:/var/www/
```
# Запуск в фоновом режиме
nohup python run.py > logs/out.log 2>&1 &
Использование с Nginx
```nginx
location /api/ {
proxy_pass http://localhost:8000;
proxy_set_header Host $host;
proxy_set_header X-Real-IP $remote_addr;
}
```
## 🤝 Вклад в проект
- Fork репозитория
- Создайте ветку для вашей фичи (git checkout -b feature/amazing-feature)
- Commit изменений (git commit -m 'Add amazing feature')
- Push в ветку (git push origin feature/amazing-feature)
- Откройте Pull Request
**📄 Лицензия**
MIT License
👥 Команда
Грачев Денис — Разработка - GitHub
### 🙏 Благодарности
***Oxford-IIIT Pet Dataset*** для обучения модели
***Сообществу PyTorch и FastAPI***
📞 Контакты
- 📧 Email: your.email@example.com
- 🐦 Telegram: @oxydencher
- 🐙 GitHub: gdg6
<div align="center"> <sub>Built with ❤️ for the Bone Quality Assessment Hackathon</sub> </div>

BIN
example.jpg Executable file

Binary file not shown.

After

Width:  |  Height:  |  Size: 32 KiB

BIN
example2.jpg Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 494 KiB

137
inference.py Normal file
View File

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

6
main.py Normal file
View File

@ -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/

BIN
prediction_result.png Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 538 KiB

BIN
public/static/arch.png Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 458 KiB

19
pydicom_custom.py Normal file
View File

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

47
requirements.txt Normal file
View File

@ -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

31
src/__init__.py Normal file
View File

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

10
src/api/__init__.py Normal file
View File

@ -0,0 +1,10 @@
"""
REST API для сервиса оценки качества
"""
from src.api.endpoints import app
# Можно добавить middleware или настройки
__all__ = [
"app",
]

267
src/api/endpoints.py Normal file
View File

@ -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"
)

875
src/api/static/index.html Normal file
View File

@ -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>

9
src/config.py Normal file
View File

@ -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 на хакатоне

49
src/data_loader.py Normal file
View File

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

View File

View File

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

41
src/model/__init__.py Normal file
View File

@ -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}")

198
src/model/segmentator.py Normal file
View File

@ -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

122
src/model/unet.py Normal file
View File

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

23
src/quality/__init__.py Normal file
View File

@ -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",
]

View File

@ -0,0 +1,3 @@
class ArtifactDetector():
def __int__(self):
pass

View File

@ -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}

View File

@ -0,0 +1,3 @@
class PositionValidator():
def __init__(self):
pass

View File

@ -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

View File

@ -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}")

285
train.py Normal file
View File

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

BIN
training_history.png Normal file

Binary file not shown.

After

Width:  |  Height:  |  Size: 76 KiB