develop - hack_2026

This commit is contained in:
denis 2026-09-22 23:29:28 +03:00
parent c40cde7ec0
commit 3121ecb8a3
17 changed files with 2980 additions and 0 deletions

22
QWEN.md
View File

@ -214,6 +214,28 @@ The system automatically determines the anatomical region from the DICOM image c
---
## Multi-Model Architecture (Planned)
See `docs/multi_model_architecture.md` for the planned pipeline:
```
Pipeline:
1. Region Detector → 2. Segmentator → 3. Quality Classifier → 4. Violation Type → 5. Aggregator
```
### Planned Models:
| Model | Purpose | File |
|-------|---------|------|
| Region Detector | Spine/Hip detection | `src/models/region_detector.py` |
| Segmentator | Bone segmentation | `src/models/segmentation/` |
| Quality Classifier | OK/Violation binary | `src/models/classification/quality.py` |
| Violation Classifier | 7+ violation types | `src/models/classification/violation.py` |
| Artifact Detector | Motion, metal detection | `src/models/artifacts/detector.py` |
| Grad-CAM | Attention heatmap | `src/models/visualization/gradcam.py` |
---
## Docker
```bash

View File

@ -0,0 +1,547 @@
# Архитектура мультимодельного пайплайна
## Обзор
Предлагается архитектура из нескольких специализированных моделей, объединённых в единый пайплайн обработки.
```
┌─────────────────────────────────────────────────────────────────────────────┐
│ DXA Quality Assessment Pipeline │
├─────────────────────────────────────────────────────────────────────────────┤
│ │
│ ┌──────────┐ ┌──────────────┐ ┌────────────────────┐ │
│ │ Input │───▶│ Region │───▶│ Segmentation │ │
│ │ DICOM │ │ Detector │ │ (ROI Model) │ │
│ └──────────┘ └──────────────┘ └────────┬─────────┘ │
│ │ │
│ ▼ │
│ ┌──────────────────────────────────────────────────────────────────┐ │
│ │ Quality Assessment │ │
│ │ ┌────────────────┐ ┌────────────────┐ ┌────────────────┐ │ │
│ │ │ Quality │ │ Violation │ │ Artifact │ │ │
│ │ │ Classifier │ │ Type │ │ Detector │ │ │
│ │ │ (OK/Violation)│ │ Classifier │ │ (motion, │ │ │
│ │ │ │ │ (7+ types) │ │ metal) │ │ │
│ │ └────────────────┘ └────────────────┘ └────────────────┘ │ │
│ └──────────────────────────────────────────────────────────────────┘ │
│ │ │
│ ▼ │
│ ┌──────────────────────────────────────────────────────────────────┐ │
│ │ Result Aggregator │ │
│ │ - quality_class (0/1) │ │
│ │ - violation_type (detailed) │ │
│ │ - reason (human-readable) │ │
│ │ - confidence_per_class │ │
│ │ - visualization (mask, heatmap) │ │
│ └──────────────────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────────────────┘
```
---
## Модели
### 1. Region Detector (Определение анатомической области)
**Назначение**: Определение типа исследования (spine/hip_left/hip_right)
**Архитектура**:
- EfficientNet-B0 или ResNet18
- 4 класса вывода: spine, hip_left, hip_right, hip
- Вход: 224x224 RGB
**Файл**: `src/models/region_detector.py`
```python
class RegionDetector(nn.Module):
"""Определение анатомической области"""
def __init__(self, num_classes=4):
super().__init__()
self.backbone = models.efficientnet_b0(weights='IMAGENET1K_V1')
self.backbone.classifier = nn.Linear(1280, num_classes)
def forward(self, x):
return self.backbone(x)
```
**Обучение**:
```bash
python src/models/train_region_detector.py --epochs 20
```
---
### 2. Segmentation Model (Сегментация ROI)
**Назначение**: Сегментация костных структур для дальнейшего анализа
**Варианты**:
#### Вариант A: Универсальная модель (рекомендуемый)
| Регион | Выход |
|--------|--------|
| Spine | Маска позвонков L1-L4 |
| Hip | Маска бедренной кости + шейка + вертелы |
**Архитектура**: U-Net с ResNet34 encoder
```python
class DXASegmentator(nn.Module):
"""Сегментация костных структур"""
def __init__(self, in_channels=3, out_classes=2):
super().__init__()
self.encoder = models.resnet34(pretrained=True)
self.decoder = UNetDecoder(512, out_classes)
def forward(self, x):
features = self.encoder(x)
return self.decoder(features)
```
**Классы сегментации**:
- Background (0)
- Bone/Vertebrae (1)
- ROI specific: Neck, Trochanter (2) - для hip
#### Вариант B: Отдельные модели
| Модель | Назначение | Выход |
|--------|------------|--------|
| SpineSegmentator | Сегментация позвонков | L1-L4 маски |
| HipSegmentator | Сегментация бедра | Femur + Neck + Trochanters |
| ROISegmenter | Сегментация ROI | Области измерения |
**Файл**: `src/models/segmentation/`
```
src/models/segmentation/
├── __init__.py
├── base.py # Базовый класс
├── spine_model.py # Сегментация позвоночника
├── hip_model.py # Сегментация бедра
└── roi_model.py # Сегментация ROI
```
---
### 3. Quality Classifier (Классификатор качества)
**Назначение**: Бинарная классификация - пригодно/непригодно
**Архитектура**: ResNet18 + attention mechanism
```python
class QualityClassifier(nn.Module):
"""
Бинарный классификатор качества
Вход: изображение + маска сегментации
"""
def __init__(self, backbone='resnet18'):
super().__init__()
self.image_encoder = models.resnet18(weights='IMAGENET1K_V1')
self.image_encoder.fc = nn.Identity()
# Дополнительный вход для маски
self.mask_encoder = nn.Sequential(
nn.Conv2d(1, 32, 3, padding=1),
nn.ReLU(),
nn.Conv2d(32, 64, 3, stride=2),
nn.ReLU(),
nn.AdaptiveAvgPool2d(1)
)
# Объединение признаков
self.fusion = nn.Linear(512 + 64, 256)
self.classifier = nn.Linear(256, 2) # OK / Violation
def forward(self, image, mask):
img_features = self.image_encoder(image)
mask_features = self.mask_encoder(mask).squeeze()
combined = torch.cat([img_features, mask_features], dim=1)
return self.classifier(self.fusion(combined))
```
**Выход**:
- quality_class: 0 (OK) или 1 (Violation)
- confidence: вероятность
---
### 4. Violation Type Classifier (Классификатор типа нарушения)
**Назначение**: Определение конкретного типа нарушения
**Архитектура**: Multi-task learning с shared backbone
```python
class ViolationTypeClassifier(nn.Module):
"""
Многоклассовый классификатор типа нарушения
"""
# Типы нарушений для spine
SPINE_VIOLATIONS = [
'correct',
'position_error', # Неправильная укладка
'artifact_motion', # Артефакт движения
'artifact_other', # Другие артефакты
'labeling_error', # Ошибка разметки
'incomplete_view', # Неполный вид
'roi_error' # Ошибка ROI
]
# Типы нарушений для hip
HIP_VIOLATIONS = [
'correct',
'position_error',
'rotation', # Нарушение ротации
'artifact_motion',
'artifact_other',
'roi_error',
'incomplete_view'
]
def __init__(self, backbone='resnet18', num_classes=7):
super().__init__()
self.backbone = models.resnet18(weights='IMAGENET1K_V1')
self.backbone.fc = nn.Identity()
# Голова для каждого типа нарушения
self.heads = nn.ModuleDict({
'spine': nn.Linear(512, 7),
'hip': nn.Linear(512, 7)
})
self.feature_dim = 512
def forward(self, x, region_type='spine'):
features = self.backbone(x)
return self.heads[region_type](features)
```
**Обучение с Weighted Loss**:
```python
# Weighted CrossEntropy для дисбаланса классов
class_weights = torch.tensor([1.0, 3.0, 2.5, 2.0, 2.5, 3.0, 2.0])
criterion = nn.CrossEntropyLoss(weight=class_weights)
```
---
### 5. Artifact Detector (Детектор артефактов)
**Назначение**: Специализированная детекция артефактов
```python
class ArtifactDetector(nn.Module):
"""
Детектор артефактов - специализированная модель
"""
# Типы артефактов
ARTIFACT_TYPES = [
'motion_blur', # Размытие движения
'metal', # Металлические объекты
'implant', # Имплантаты/протезы
'cement', # Цемент
'calcification', # Кальцинаты
'noise', # Шум
'clipping' # Обрезанное изображение
]
def __init__(self):
super().__init__()
# EfficientNet-B0 для извлечения признаков
self.encoder = models.efficientnet_b0(weights='IMAGENET1K_V1')
self.encoder.classifier = nn.Identity()
# Голова для детекции артефактов (Binary для каждого типа)
self.artifact_heads = nn.ModuleList([
nn.Linear(1280, 2) # есть/нет для каждого типа
for _ in self.ARTIFACT_TYPES
])
def forward(self, x):
features = self.encoder(x)
outputs = [head(features) for head in self.artifact_heads]
return torch.stack(outputs, dim=1) # [B, 7, 2]
```
**Альтернатива**: Использовать rule-based методы + простые эвристики (текущая реализация в `detailed_assessment.py`)
---
### 6. Attention/Heatmap Model (Модель внимания)
**Назначение**: Визуализация областей, на которые обращает внимание модель
```python
class GradCAMExtractor:
"""Извлечение attention heatmap с помощью Grad-CAM"""
def __init__(self, model, target_layer):
self.model = model
self.target_layer = target_layer
self.gradients = None
def save_gradient(self, grad):
self.gradients = grad
def get_heatmap(self, input_tensor, target_class):
# Forward
output = self.model(input_tensor)
# Backward
self.model.zero_grad()
class_loss = output[0, target_class]
class_loss.backward()
# Get gradient
gradients = self.gradients
pooled_gradients = torch.mean(gradients, dim=[0, 2, 3])
# Apply to features
features = self.target_layer.output
for i in range(features.shape[1]):
features[:, i, :, :] *= pooled_gradients[i]
# Create heatmap
heatmap = torch.mean(features, dim=1).squeeze()
heatmap = F.relu(heatmap)
heatmap /= torch.max(heatmap)
return heatmap
```
---
## Пайплайн обработки
### Основной класс
```python
class QualityAssessmentPipeline:
"""
Основной пайплайн оценки качества DXA исследований
"""
def __init__(self, config: PipelineConfig):
self.config = config
self.device = get_device()
# Загрузка моделей
self.region_detector = self._load_model('region_detector')
self.segmentator = self._load_model('segmentator')
self.quality_classifier = self._load_model('quality_classifier')
self.violation_classifier = self._load_model('violation_classifier')
self.artifact_detector = self._load_model('artifact_detector') # опционально
# Правила для финального решения
self.rules = QualityRules()
def analyze(self, dicom_path: str) -> QualityResult:
"""Полный анализ DICOM файла"""
# 1. Загрузка и предобработка
image, ds = self._load_dicom(dicom_path)
# 2. Определение региона
region = self.region_detector.predict(image)
# 3. Сегментация
segmentation = self.segmentator.predict(image, region)
# 4. Классификация качества
quality_result = self.quality_classifier.predict(image, segmentation)
# 5. Определение типа нарушения (если есть)
if quality_result.predicted_class == 1:
violation_type = self.violation_classifier.predict(
image, segmentation, region
)
else:
violation_type = 'correct'
# 6. Детекция артефактов (дополнительно)
artifacts = self.artifact_detector.predict(image) if self.artifact_detector else {}
# 7. Применение правил
final_result = self.rules.apply(
quality_result=quality_result,
violation_type=violation_type,
artifacts=artifacts,
segmentation=segmentation,
region=region
)
# 8. Визуализация
heatmap = self._generate_heatmap(image, quality_result)
return QualityResult(
**final_result,
segmentation=segmentation,
heatmap=heatmap
)
def _load_dicom(self, path):
"""Загрузка DICOM"""
ds = pydicom.dcmread(path)
image = ds.pixel_array.astype(np.float32)
image = (image - image.min()) / (image.max() - image.min() + 1e-8)
return image, ds
```
---
## Конфигурация
```python
@dataclass
class PipelineConfig:
"""Конфигурация пайплайна"""
# Модели
region_detector_path: str = "models/region_detector.pth"
segmentator_path: str = "models/segmentator.pth"
quality_classifier_path: str = "models/quality_classifier.pth"
violation_classifier_path: str = "models/violation_classifier.pth"
artifact_detector_path: Optional[str] = None
# Параметры
input_size: int = 224
device: str = "auto" # auto, cuda, mps, cpu
# Режимы
use_heatmap: bool = True
use_artifact_detector: bool = False
use_multi_gpu: bool = False
# Threshold
quality_threshold: float = 0.5
confidence_threshold: float = 0.7
```
---
## Файловая структура
```
src/
├── models/
│ ├── __init__.py
│ ├── base.py # Базовые классы
│ ├── registry.py # Регистр моделей
│ ├── factory.py # Фабрика моделей
│ │
│ ├── region_detector.py # Определение региона
│ ├── segmentator.py # Сегментация (общая)
│ │
│ ├── classification/
│ │ ├── __init__.py
│ │ ├── quality_classifier.py # OK/Violation
│ │ └── violation_classifier.py # Типы нарушений
│ │
│ ├── segmentation/
│ │ ├── __init__.py
│ │ ├── spine_segmentator.py # Позвоночник
│ │ └── hip_segmentator.py # Бедро
│ │
│ ├── artifacts/
│ │ ├── __init__.py
│ │ └── detector.py # Детектор артефактов
│ │
│ └── visualization/
│ ├── __init__.py
│ └── gradcam.py # Attention heatmap
│
├── pipeline/
│ ├── __init__.py
│ ├── config.py # Конфигурация
│ ├── pipeline.py # Основной пайплайн
│ ├── rules.py # Правила принятия решений
│ └── result.py # Результат
│
├── training/
│ ├── __init__.py
│ ├── train_region.py
│ ├── train_segmentation.py
│ ├── train_quality.py
│ ├── train_violation.py
│ └── losses.py # Weighted losses
│
└── evaluation/
├── __init__.py
├── metrics.py # ROC-AUC, PR-AUC, F1
└── analysis.py # Анализ ошибок
```
---
## Обучение
### Скрипт обучения
```bash
# 1. Обучение region detector
python src/training/train_region.py \
--data-root dataset_hack \
--annotation dataset_hack/разметка.xlsx \
--epochs 20 \
--batch-size 16 \
--output models/region_detector.pth
# 2. Обучение сегментатора
python src/training/train_segmentation.py \
--data-root dataset_hack \
--epochs 30 \
--batch-size 8 \
--output models/segmentator.pth
# 3. Обучение quality classifier
python src/training/train_quality.py \
--data-root dataset_hack \
--annotation dataset_hack/разметка.xlsx \
--epochs 25 \
--batch-size 16 \
--loss weighted \
--output models/quality_classifier.pth
# 4. Обучение violation classifier
python src/training/train_violation.py \
--data-root dataset_hack \
--annotation dataset_hack/разметка.xlsx \
--epochs 25 \
--batch-size 16 \
--num-classes 7 \
--loss weighted \
--output models/violation_classifier.pth
```
---
## Метрики
| Модель | Метрика | Целевое значение |
|--------|---------|-----------------|
| Region Detector | Accuracy | > 95% |
| Segmentator | Dice | > 0.85 |
| Quality Classifier | ROC-AUC | > 0.90 |
| Quality Classifier | F1 | > 0.85 |
| Violation Classifier | Macro-F1 | > 0.75 |
| Violation Classifier | PR-AUC | > 0.80 |
---
## Приоритеты реализации
### Фаза 1: Основные модели
1. ✅ Region Detector (уже есть в inference.py)
2. Quality Classifier (улучшить)
3. Violation Type Classifier (новый)
### Фаза 2: Сегментация
4. Spine Segmentator
5. Hip Segmentator
### Фаза 3: Улучшения
6. Artifact Detector
7. Grad-CAM Heatmap
8. Multi-task learning

393
docs/presentation.md Normal file
View File

@ -0,0 +1,393 @@
# DXA Quality Assessment — Документация для презентации
## Содержание
1. [Обзор проекта](#1-обзор-проекта)
2. [Проблема и актуальность](#2-проблема-и-актуальность)
3. [Техническая архитектура](#3-техническая-архитектура)
4. [Датасет и разметка](#4-датасет-и-разметка)
5. [ML-модель](#5-ml-модель)
6. [API и интеграция](#6-api-и-интеграция)
7. [Результаты](#7-результаты)
8. [Демонстрация](#8-демонстрация)
9. [Перспективы развития](#9-перспективы-развития)
---
## 1. Обзор проекта
### Что это за проект?
**DXA Quality Assessment** — это AI-сервис для автоматизированной оценки качества денситометрических исследований (DXA — Dual-energy X-ray Absorptiometry).
### Ключевые возможности
| Возможность | Описание |
|-------------|----------|
| 📥 Загрузка DICOM | Работа с медицинскими изображениями в стандартном формате |
| 🔍 Анализ качества | Оценка по критериям: позиционирование, артефакты, контраст |
| 🏷️ Классификация | Бинарная классификация: OK / Нарушение |
| 📊 Экспорт | Вывод результатов в XLSX/CSV формате |
| 🌐 REST API | Интеграция с внешними системами |
| 🖥️ Web-интерфейс | Удобная визуализация и загрузка |
### Типы анализируемых исследований
- **Позвоночник** (spine)
- **Бедро правое** (hip_right)
- **Бедро левое** (hip_left)
---
## 2. Проблема и актуальность
### Проблема
При проведении денситометрии (DXA) критически важно соблюдать стандарты качества исследования:
- ❌ Неправильное позиционирование пациента
- ❌ Артефакты движения
- ❌ Недостаточный контраст
- ❌ Неправильная укладка
Нарушение этих критериев приводит к:
- ❌ Неправильной диагностике остеопороза
- ❌ Повторным исследованиям
- ❌ Дополнительной лучевой нагрузке
- ❌ Потере времени и ресурсов
### Решение
Автоматизированная система оценки качества, которая:
- ✅ Мгновенно анализирует каждое изображение
- ✅ Выявляет конкретные нарушения
- ✅ Обеспечивает стандартизацию контроля качества
- ✅ Интегрируется в рабочий процесс рентгенолога
### Бизнес-ценность
```
Экономия времени рентгенолога: ~30-60 секунд на исследование
Снижение количества повторных исследований: до 15-20%
Повышение качества диагностики: стандартизированный контроль
```
---
## 3. Техническая архитектура
### Общая схема
```
┌─────────────┐ ┌──────────────┐ ┌─────────────┐
│ DICOM │────▶│ FastAPI │────▶│ PyTorch │
│ Файлы │ │ Server │ │ Model │
└─────────────┘ └──────────────┘ └─────────────┘
│
▼
┌──────────────┐
│ XLSX/CSV │
│ Output │
└──────────────┘
```
### Компоненты системы
| Компонент | Технология | Назначение |
|-----------|------------|------------|
| Backend | Python 3.10, FastAPI | REST API, веб-сервер |
| ML Framework | PyTorch 2.1 | Глубокое обучение |
| Модель | ResNet18 (ImageNet) | Классификация изображений |
| Обработка изображений | PIL, OpenCV | Предобработка DICOM |
| Данные | pandas, openpyxl | Работа с Excel |
| Контейнеризация | Docker | Деплой |
### Структура проекта
```
bone_2026/
├── src/
│ ├── main.py # FastAPI приложение
│ ├── dxa/
│ │ ├── dataset.py # Загрузка данных
│ │ ├── model.py # Архитектура модели
│ │ ├── train.py # Обучение
│ │ └── inference.py # Инференс
│ └── api/ # Эндпоинты
├── models/
│ └── dxa_model.pth # Обученная модель
├── dataset_hack/ # Датасет
└── requirements.txt # Зависимости
```
---
## 4. Датасет и разметка
### Источник данных
- **100 исследований** из реальной клинической практики
- **499 DICOM файлов**
- **Аннотации** в формате Excel
### Типы анатомических регионов
| Регион | Критерии оценки |
|--------|-----------------|
| Позвоночник | Укладка, Ось, Артефакты |
| Бедро (правое/левое) | Позиция, ROI |
### Классы качества
| Класс | Описание | Примерное соотношение |
|-------|----------|----------------------|
| 0 | OK (качество соответствует стандартам) | ~65% |
| 1 | Нарушение (выявлены проблемы) | ~35% |
### Пример разметки
```
Study UID | Позвоночник_итого | Бедро_прав_итого | Бедро_лев_итого | Класс
----------|-------------------|------------------|-----------------|------
1.2.840...| 0 | 1 | 0 | 1
1.2.840...| 0 | 0 | 0 | 0
```
### Разделение данных
- **Обучение**: ~80% (1146 samples)
- **Валидация**: ~20% (287 samples)
---
## 5. ML-модель
### Архитектура
```
DXAQualityClassifier (ResNet18)
├── Backbone: ResNet18 (pretrained on ImageNet)
│ └── Feature extraction: 512 dimensions
└── Classifier Head:
├── Dropout (0.3)
├── Linear (512 → 256)
├── ReLU
├── Dropout (0.3)
└── Linear (256 → 2)
```
### Входные данные
- **Размер изображения**: 224 × 224 пикселя
- **Формат**: RGB (3 канала)
- **Нормализация**: ImageNet stats
### Выходные данные
- **Класс 0**: Качество OK
- **Класс 1**: Нарушение качества
- **Confidence**: Вероятность предсказания
### Гиперпараметры обучения
| Параметр | Значение |
|----------|----------|
| Backbone | ResNet18 |
| Input size | 224 × 224 |
| Batch size | 8-16 |
| Learning rate | 1e-4 |
| Optimizer | AdamW |
| Scheduler | ReduceLROnPlateau |
| Dropout | 0.3 |
### Пример команды обучения
```bash
python src/dxa/train.py \
--epochs 10 \
--batch-size 16 \
--backbone resnet18 \
--data-root dataset_hack
```
---
## 6. API и интеграция
### REST Endpoints
| Метод | Эндпоинт | Описание |
|-------|----------|----------|
| GET | `/` | Главная страница (web UI) |
| GET | `/api/v1/health` | Проверка статуса |
| POST | `/api/v1/analyze` | Анализ одного DICOM |
| POST | `/api/v1/batch` | Пакетный анализ |
| POST | `/api/v1/export` | Анализ + экспорт XLSX |
### Пример запроса (cURL)
```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.dcm"
```
### Пример ответа
```json
{
"study_uid": "1.2.840.113619.2.110.512719.20250403090444",
"image_uid": "1.2.3.4.5.6.7.8.9",
"anatomical_region": "spine",
"quality_class": 0,
"quality_label": "OK",
"confidence": 0.9234,
"processing_status": "Success"
}
```
### Формат вывода (XLSX/CSV)
| Колонка | Описание |
|---------|----------|
| path_to_study | Путь к исследованию |
| study_uid | StudyInstanceUID |
| image_uid | SOPInstanceUID |
| anatomical_region | Регион (spine/hip) |
| quality_class | Класс (0/1) |
| violation_type | Тип нарушения |
| processing_status | Статус обработки |
| time_of_processing | Время обработки (сек) |
### Web-интерфейс
- 📤 Drag-and-drop загрузка
- 🔍 Автоматический анализ
- 📊 Визуализация результатов
- 📥 Экспорт в Excel
---
## 7. Результаты
### Метрики на валидации
| Метрика | Значение |
|---------|----------|
| Accuracy | ~84% |
| Precision | ~0.35 |
| Recall | ~0.22 |
| **F1 Score** | **~0.27** |
### Анализ результатов
```
Распределение классов в валидации:
├── OK (класс 0): ████████████████████ 65%
└── Нарушение (класс 1): ██████████ 35%
Примечание: F1 скорее всего ниже из-за
несбалансированности классов и ограниченного
количества данных (100 исследований)
```
### Время обработки
- **Одно изображение**: ~50-100 мс
- **Пакетная обработка**: зависит от размера батча
### Ограничения текущей версии
- ⚠️ Ограниченный объём данных (100 исследований)
- ⚠️ Дисбаланс классов
- ⚠️ Только бинарная классификация
- ⚠️ Требуется больше эпох обучения
### Рекомендации для улучшения
1. **Увеличение датасета**: 500+ исследований
2. **Балансировка классов**: weighted loss, oversampling
3. **Аугментация**: rotation, flip, brightness
4. **Архитектура**: EfficientNet, ViT
5. **Многоклассовая классификация**: типы нарушений
---
## 8. Демонстрация
### Сценарий использования
1. **Врач загружает DICOM файл** через web-интерфейс или API
2. **Система автоматически:**
- Определяет анатомический регион
- Анализирует качество изображения
- Классифицирует: OK / Нарушение
3. **Результат:**
- Визуализация в интерфейсе
- Excel-отчёт для дальнейшего анализа
### Пример работы
```
Вход: DICOM файл исследования позвоночника
↓
Обработка: ResNet18 классификатор
↓
Выход: {
"anatomical_region": "spine",
"quality_class": 0,
"confidence": 0.89
}
```
### Screenshots (доступны в интерфейсе)
- Загрузка файла
- Результат анализа
- Экспорт отчёта
---
## 9. Перспективы развития
### Краткосрочные улучшения
| Направление | Описание |
|-------------|----------|
| Увеличение данных | Добавить 400+ исследований |
| Классовый баланс | Weighted loss, SMOTE |
| Аугментация | Rotation, flip, brightness |
| Типы нарушений | Multiclass: позиция, артефакты, контраст |
### Среднесрочные цели
| Направление | Описание |
|-------------|----------|
| Сегментация | Выделение позвонков, бедра |
| Многозадачность | Классификация + сегментация |
| AutoML | Оптимальная архитектура |
| Мониторинг | Логирование, метрики |
### Долгосрочное видение
- 🏥 **Интеграция с PACS**
- 📱 **Мобильное приложение**
- 🤖 **Real-time анализ**
- 🌐 **Облачное решение (SaaS)**
- 📊 **Дашборд для радиологов**
---
## Контакты
| | |
|---|---|
| **Автор** | Грачев Денис |
| **Email** | denis@example.com |
| **Telegram** | @oxydencher |
| **GitHub** | github.com/yourusername |
---
*Проект разработан в рамках хакатона по медицинскому ИИ*

111
docs/tasks_next_session.md Normal file
View File

@ -0,0 +1,111 @@
# Задачи для следующей сессии
## Сравнение текущей реализации с требованиями
### ✅ РЕАЛИЗОВАНО
| Требование | Статус | Файл |
|------------|--------|------|
| Бинарная классификация OK/Нарушение | ✅ Готово | src/dxa/model.py |
| Определение анатомической области | ✅ Готово | src/dxa/inference.py |
| Экспорт в XLSX/CSV | ✅ Готово | src/main.py /api/v1/export |
| API для пакетной обработки | ✅ Готово | src/main.py /api/v1/batch |
| Docker контейнеризация | ✅ Готово | Dockerfile |
| Web-интерфейс | ✅ Готово | src/api/static/ |
| Определение типа нарушения | ✅ Частично | src/quality/detailed_assessment.py |
| Детальные метрики качества | ✅ Готово | src/quality/detailed_assessment.py |
| Проверка движения (motion) | ✅ Готово | detect_motion_blur() |
| Проверка артефактов | ✅ Готово | detect_artifacts() |
| DICOM SR формат | ✅ Готово | /api/v1/analyze/sr |
| Визуализация маски | ✅ Готово | include_visualization=true |
---
## Мультимодельная архитектура
См. документацию: `docs/multi_model_architecture.md`
```
Pipeline:
1. Region Detector → 2. Segmentator → 3. Quality Classifier → 4. Violation Type → 5. Aggregator
```
### Модели:
| Модель | Назначение | Файл |
|--------|------------|------|
| Region Detector | Spine/Hip определение | `src/models/region_detector.py` |
| Segmentator | Сегментация костей | `src/models/segmentation/` |
| Quality Classifier | OK/Violation бинарный | `src/models/classification/quality.py` |
| Violation Classifier | Типы нарушений (7+) | `src/models/classification/violation.py` |
| Artifact Detector | Детекция артефактов | `src/models/artifacts/detector.py` |
| Grad-CAM | Heatmap внимания | `src/models/visualization/gradcam.py` |
---
## ПЛАН РЕАЛИЗАЦИИ
### Фаза 1: Архитектура и модели (высокий приоритет)
1. **Создать структуру моделей**
- `src/models/` - директория для всех моделей
- `src/pipeline/` - основной пайплайн
- `src/training/` - скрипты обучения
2. **Region Detector**
- Переиспользовать текущую логику из inference.py
- Обучить классификатор на 4 класса
3. **Segmentation Models**
- Spine Segmentator (U-Net)
- Hip Segmentator (U-Net)
4. **Quality Classifier (бинарный)**
- ResNet18 с attention
- Weighted loss для балансировки
5. **Violation Type Classifier (многоклассовый)**
- 7 классов для spine
- 7 классов для hip
- Multi-task learning
### Фаза 2: Пайплайн и интеграция
6. **Pipeline класс**
- Объединение всех моделей
- Queue для асинхронной обработки
7. **API интеграция**
- Обновить endpoints для использования пайплайна
- Batch processing
### Фаза 3: Визуализация и метрики
8. **Grad-CAM Heatmap**
- Извлечение attention карт
- Наложение на изображение
9. **Метрики**
- ROC-AUC, PR-AUC
- Macro-F1 с доверительными интервалами
---
## Текущие ограничения
1. Модель обучена на бинарную классификацию - нужно переобучение
2. F1 ~0.27 - низкий из-за дисбаланса классов
3. Датасет ~100 исследований - нужно 500+
4. Сегментация использует простой threshold - нужно дообучить модель
5. Heatmap не реализовано - требует дообучения
---
## ЗАПУСК ДЛЯ ТЕСТИРОВАНИЯ
```bash
cd /Users/denis/workspace/bone_2026
python -m uvicorn src.main:app --host 0.0.0.0 --port 8000
```
Открыть http://localhost:8000

32
src/models/__init__.py Normal file
View File

@ -0,0 +1,32 @@
"""
DXA Quality Assessment Models
"""
from .base import BaseDXAModel, BaseClassifier, BaseSegmentator
from .region_detector import RegionDetector, create_region_detector, determine_region_from_image
from .classification import QualityClassifier, ViolationTypeClassifier
from .segmentation import DXASegmenter, create_segmenter
from .visualization import GradCAMExtractor, visualize_attention
__all__ = [
# Base
'BaseDXAModel',
'BaseClassifier',
'BaseSegmentator',
# Region
'RegionDetector',
'create_region_detector',
'determine_region_from_image',
# Classification
'QualityClassifier',
'ViolationTypeClassifier',
# Segmentation
'DXASegmenter',
'create_segmenter',
# Visualization
'GradCAMExtractor',
'visualize_attention'
]

92
src/models/base.py Normal file
View File

@ -0,0 +1,92 @@
"""
Base model class for all DXA models
"""
import torch
import torch.nn as nn
from typing import Dict, Any, Optional
from abc import ABC, abstractmethod
class BaseDXAModel(nn.Module, ABC):
"""Base class for all DXA models"""
def __init__(self, config: Optional[Dict] = None):
super().__init__()
self.config = config or {}
self.device = torch.device('cpu')
@abstractmethod
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""Forward pass"""
pass
def extract_features(self, x: torch.Tensor) -> torch.Tensor:
"""Extract features without classification"""
return self.forward(x)
def to_device(self, device: str = 'cpu'):
"""Move model to device"""
self.device = torch.device(device)
self.to(self.device)
return self
def get_info(self) -> Dict[str, Any]:
"""Get model info"""
return {
'name': self.__class__.__name__,
'device': str(self.device),
'config': self.config
}
class BaseClassifier(BaseDXAModel):
"""Base classifier"""
def __init__(self, num_classes: int = 2, **kwargs):
super().__init__(**kwargs)
self.num_classes = num_classes
@abstractmethod
def predict(self, x: torch.Tensor) -> Dict[str, Any]:
"""Predict on input"""
pass
def predict_proba(self, x: torch.Tensor) -> torch.Tensor:
"""Get probabilities"""
self.eval()
with torch.no_grad():
logits = self.forward(x)
return torch.softmax(logits, dim=1)
def predict_class(self, x: torch.Tensor) -> torch.Tensor:
"""Get class predictions"""
probs = self.predict_proba(x)
return probs.argmax(dim=1)
class BaseSegmentator(BaseDXAModel):
"""Base segmentator"""
def __init__(self, in_channels: int = 3, out_classes: int = 2, **kwargs):
super().__init__(**kwargs)
self.in_channels = in_channels
self.out_classes = out_classes
@abstractmethod
def predict(self, x: torch.Tensor) -> Dict[str, Any]:
"""Predict on input"""
pass
def predict_mask(self, x: torch.Tensor) -> torch.Tensor:
"""Get segmentation mask"""
self.eval()
with torch.no_grad():
logits = self.forward(x)
return logits.argmax(dim=1)
def predict_proba(self, x: torch.Tensor) -> torch.Tensor:
"""Get class probabilities"""
self.eval()
with torch.no_grad():
logits = self.forward(x)
return torch.softmax(logits, dim=1)

View File

@ -0,0 +1,21 @@
"""
Classification models
"""
from .quality_classifier import QualityClassifier, create_quality_classifier
from .violation_classifier import (
ViolationTypeClassifier,
create_violation_classifier,
SPINE_VIOLATIONS,
HIP_VIOLATIONS,
VIOLATION_DESCRIPTIONS
)
__all__ = [
'QualityClassifier',
'create_quality_classifier',
'ViolationTypeClassifier',
'create_violation_classifier',
'SPINE_VIOLATIONS',
'HIP_VIOLATIONS',
'VIOLATION_DESCRIPTIONS'
]

View File

@ -0,0 +1,224 @@
"""
Quality Classifier - бинарный классификатор OK/Violation
"""
import torch
import torch.nn as nn
import torchvision.models as models
from typing import Dict, Tuple, Optional
import numpy as np
class QualityClassifier(nn.Module):
"""
Бинарный классификатор качества DXA исследования.
Выход: 0 (OK) или 1 (Violation)
"""
QUALITY_LABELS = {0: 'OK', 1: 'Violation'}
def __init__(self,
backbone: str = 'resnet18',
pretrained: bool = True,
dropout: float = 0.4,
use_attention: bool = False):
super().__init__()
self.backbone_name = backbone
self.use_attention = use_attention
# Encoder для изображения
if backbone == 'resnet18':
self.image_encoder = models.resnet18(
weights='IMAGENET1K_V1' if pretrained else None
)
self.image_encoder.fc = nn.Identity()
image_features = 512
elif backbone == 'resnet34':
self.image_encoder = models.resnet34(
weights='IMAGENET1K_V1' if pretrained else None
)
self.image_encoder.fc = nn.Identity()
image_features = 512
elif backbone == 'efficientnet_b0':
self.image_encoder = models.efficientnet_b0(
weights='IMAGENET1K_V1' if pretrained else None
)
self.image_encoder.classifier = nn.Identity()
image_features = 1280
else:
raise ValueError(f"Unknown backbone: {backbone}")
# Encoder для маски (опционально)
if use_attention:
self.mask_encoder = nn.Sequential(
nn.Conv2d(1, 32, 3, padding=1),
nn.BatchNorm2d(32),
nn.ReLU(inplace=True),
nn.Conv2d(32, 64, 3, stride=2, padding=1),
nn.BatchNorm2d(64),
nn.ReLU(inplace=True),
nn.AdaptiveAvgPool2d(1)
)
combined_features = image_features + 64
else:
self.mask_encoder = None
combined_features = image_features
# Классификатор
self.classifier = nn.Sequential(
nn.Dropout(dropout),
nn.Linear(combined_features, 256),
nn.ReLU(inplace=True),
nn.Dropout(dropout),
nn.Linear(256, 2) # OK / Violation
)
self.image_features = image_features
def forward(self, image: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
"""
Forward pass
Args:
image: [B, 3, H, W] RGB изображение
mask: [B, 1, H, W] опциональная маска сегментации
"""
# Извлечение признаков из изображения
img_features = self.image_encoder(image)
# Если есть маска и используется attention
if mask is not None and self.use_attention and self.mask_encoder is not None:
mask_features = self.mask_encoder(mask).squeeze(-1).squeeze(-1)
combined = torch.cat([img_features, mask_features], dim=1)
else:
combined = img_features
return self.classifier(combined)
def predict(self,
image: torch.Tensor,
mask: Optional[torch.Tensor] = None) -> Dict:
"""
Предсказание качества.
Returns:
Dict: predicted_class, label, confidence, probabilities
"""
self.eval()
with torch.no_grad():
logits = self.forward(image, mask)
probs = torch.softmax(logits, dim=1)
conf, pred = probs.max(dim=1)
results = []
for i in range(image.shape[0]):
pred_class = pred[i].item()
results.append({
'predicted_class': pred_class,
'label': self.QUALITY_LABELS[pred_class],
'confidence': conf[i].item(),
'probabilities': {
'OK': probs[i, 0].item(),
'Violation': probs[i, 1].item()
},
'logits': logits[i].cpu().numpy().tolist()
})
return results[0] if len(results) == 1 else results
def extract_features(self,
image: torch.Tensor,
mask: Optional[torch.Tensor] = None) -> torch.Tensor:
"""Извлечение признаков"""
img_features = self.image_encoder(image)
if mask is not None and self.use_attention and self.mask_encoder is not None:
mask_features = self.mask_encoder(mask).squeeze(-1).squeeze(-1)
return torch.cat([img_features, mask_features], dim=1)
return img_features
def get_info(self) -> Dict:
return {
'name': 'QualityClassifier',
'backbone': self.backbone_name,
'num_classes': 2,
'use_attention': self.use_attention,
'image_features': self.image_features
}
class QualityClassifierMultiRegion(nn.Module):
"""
Мультирегиональный классификатор качества.
Отдельные головы для spine и hip.
"""
def __init__(self,
backbone: str = 'resnet18',
pretrained: bool = True,
dropout: float = 0.4):
super().__init__()
# Общий encoder
if backbone == 'resnet18':
self.encoder = models.resnet18(
weights='IMAGENET1K_V1' if pretrained else None
)
self.encoder.fc = nn.Identity()
feature_dim = 512
else:
raise ValueError(f"Unknown backbone: {backbone}")
# Отдельные головы для каждого региона
self.spine_head = nn.Sequential(
nn.Dropout(dropout),
nn.Linear(feature_dim, 128),
nn.ReLU(),
nn.Linear(128, 2)
)
self.hip_head = nn.Sequential(
nn.Dropout(dropout),
nn.Linear(feature_dim, 128),
nn.ReLU(),
nn.Linear(128, 2)
)
def forward(self, x: torch.Tensor, region: str = 'spine') -> torch.Tensor:
features = self.encoder(x)
if region == 'spine':
return self.spine_head(features)
else:
return self.hip_head(features)
def predict(self, image: torch.Tensor, region: str = 'spine') -> Dict:
self.eval()
with torch.no_grad():
logits = self.forward(image, region)
probs = torch.softmax(logits, dim=1)
conf, pred = probs.max(dim=1)
return {
'predicted_class': pred.item(),
'label': QualityClassifier.QUALITY_LABELS[pred.item()],
'confidence': conf.item(),
'region': region
}
def create_quality_classifier(
backbone: str = 'resnet18',
pretrained: bool = True,
use_attention: bool = False,
device: str = 'cpu'
) -> QualityClassifier:
"""Создание классификатора качества"""
model = QualityClassifier(
backbone=backbone,
pretrained=pretrained,
use_attention=use_attention
)
model.to(device)
return model

View File

@ -0,0 +1,237 @@
"""
Violation Type Classifier - многоклассовый классификатор типа нарушения
"""
import torch
import torch.nn as nn
import torchvision.models as models
from typing import Dict, List, Optional
import numpy as np
# Типы нарушений для позвоночника
SPINE_VIOLATIONS = [
'correct',
'position_error', # Неправильная укладка
'artifact_motion', # Артефакт движения
'artifact_other', # Другие артефакты
'labeling_error', # Ошибка разметки
'incomplete_view', # Неполный вид
'roi_error' # Ошибка ROI
]
# Типы нарушений для бедра
HIP_VIOLATIONS = [
'correct',
'position_error',
'rotation', # Нарушение ротации
'artifact_motion',
'artifact_other',
'roi_error',
'incomplete_view'
]
# Русские описания
VIOLATION_DESCRIPTIONS = {
'correct': 'Качество соответствует норме',
'position_error': 'Неправильное положение пациента',
'artifact_motion': 'Артефакт движения (размытие)',
'artifact_other': 'Другие артефакты',
'labeling_error': 'Ошибка разметки',
'incomplete_view': 'Неполный вид анатомической области',
'roi_error': 'Неправильное положение ROI',
'rotation': 'Нарушение ротации'
}
class ViolationTypeClassifier(nn.Module):
"""
Многоклассовый классификатор типа нарушения.
Поддерживает разные классификаторы для spine и hip.
"""
def __init__(self,
backbone: str = 'resnet18',
num_classes_spine: int = 7,
num_classes_hip: int = 7,
pretrained: bool = True,
dropout: float = 0.4,
use_multi_task: bool = True):
super().__init__()
self.backbone_name = backbone
self.num_classes_spine = num_classes_spine
self.num_classes_hip = num_classes_hip
self.use_multi_task = use_multi_task
# Encoder
if backbone == 'resnet18':
self.encoder = models.resnet18(
weights='IMAGENET1K_V1' if pretrained else None
)
self.encoder.fc = nn.Identity()
feature_dim = 512
elif backbone == 'resnet34':
self.encoder = models.resnet34(
weights='IMAGENET1K_V1' if pretrained else None
)
self.encoder.fc = nn.Identity()
feature_dim = 512
elif backbone == 'efficientnet_b0':
self.encoder = models.efficientnet_b0(
weights='IMAGENET1K_V1' if pretrained else None
)
self.encoder.classifier = nn.Identity()
feature_dim = 1280
else:
raise ValueError(f"Unknown backbone: {backbone}")
if use_multi_task:
# Общая голова с регионом
self.head = nn.Sequential(
nn.Dropout(dropout),
nn.Linear(feature_dim, 256),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(256, 14) # 7 + 7 = 14 классов
)
else:
# Отдельные головы
self.spine_head = nn.Sequential(
nn.Dropout(dropout),
nn.Linear(feature_dim, 128),
nn.ReLU(),
nn.Linear(128, num_classes_spine)
)
self.hip_head = nn.Sequential(
nn.Dropout(dropout),
nn.Linear(feature_dim, 128),
nn.ReLU(),
nn.Linear(128, num_classes_hip)
)
self.feature_dim = feature_dim
def forward(self, x: torch.Tensor, region: str = 'spine') -> torch.Tensor:
"""
Forward pass
Args:
x: [B, 3, H, W] Input image
region: 'spine' or 'hip'
"""
features = self.encoder(x)
if self.use_multi_task:
logits = self.head(features)
# 0-6: spine, 7-13: hip
if region == 'spine':
return logits[:, :self.num_classes_spine]
else:
return logits[:, self.num_classes_spine:]
else:
if region == 'spine':
return self.spine_head(features)
else:
return self.hip_head(features)
def predict(self,
image: torch.Tensor,
region: str = 'spine',
return_all: bool = False) -> Dict:
"""
Предсказание типа нарушения.
Args:
image: Input tensor
region: 'spine' или 'hip'
return_all: вернуть все вероятности
Returns:
Dict с predicted_type, description, confidence, probabilities
"""
self.eval()
# Определение списка классов
violations = SPINE_VIOLATIONS if region == 'spine' else HIP_VIOLATIONS
with torch.no_grad():
logits = self.forward(image, region)
probs = torch.softmax(logits, dim=1)
conf, pred = probs.max(dim=1)
pred_idx = pred[0].item() if image.shape[0] == 1 else pred.item()
pred_type = violations[pred_idx]
result = {
'predicted_type': pred_type,
'description': VIOLATION_DESCRIPTIONS.get(pred_type, 'Неизвестно'),
'confidence': conf[0].item() if image.shape[0] == 1 else conf.item(),
'region': region,
'predicted_index': int(pred_idx)
}
if return_all:
result['probabilities'] = {
violations[i]: probs[0 if image.shape[0] == 1 else i, i].item()
for i in range(len(violations))
}
return result
def predict_with_aggregated(self,
image: torch.Tensor,
region: str = 'spine') -> Dict:
"""
Предсказание с агрегированными результатами для батча.
"""
self.eval()
violations = SPINE_VIOLATIONS if region == 'spine' else HIP_VIOLATIONS
with torch.no_grad():
logits = self.forward(image, region)
probs = torch.softmax(logits, dim=1)
# Средние вероятности по батчу
avg_probs = probs.mean(dim=0)
conf, pred = avg_probs.max(dim=0)
return {
'predicted_type': violations[pred.item()],
'description': VIOLATION_DESCRIPTIONS.get(violations[pred.item()], 'Неизвестно'),
'confidence': conf.item(),
'predicted_index': pred.item(),
'probabilities': {
violations[i]: avg_probs[i].item()
for i in range(len(violations))
}
}
def get_info(self) -> Dict:
return {
'name': 'ViolationTypeClassifier',
'backbone': self.backbone_name,
'num_classes_spine': self.num_classes_spine,
'num_classes_hip': self.num_classes_hip,
'spine_violations': SPINE_VIOLATIONS,
'hip_violations': HIP_VIOLATIONS,
'use_multi_task': self.use_multi_task
}
def create_violation_classifier(
backbone: str = 'resnet18',
num_classes: int = 7,
pretrained: bool = True,
device: str = 'cpu'
) -> ViolationTypeClassifier:
"""Создание классификатора"""
model = ViolationTypeClassifier(
backbone=backbone,
num_classes_spine=num_classes,
num_classes_hip=num_classes,
pretrained=pretrained
)
model.to(device)
return model

View File

@ -0,0 +1,192 @@
"""
Region Detector - определяет анатомическую область исследования
"""
import torch
import torch.nn as nn
import torchvision.models as models
from typing import Dict, List, Tuple, Optional
import numpy as np
class RegionDetector(nn.Module):
"""
Определение анатомической области исследования.
Классы: spine, hip_left, hip_right, hip (unknown side)
"""
REGIONS = ['spine', 'hip_left', 'hip_right', 'hip']
def __init__(self,
backbone: str = 'efficientnet_b0',
num_classes: int = 4,
pretrained: bool = True,
dropout: float = 0.3):
super().__init__()
self.backbone_name = backbone
self.num_classes = num_classes
# Загрузка backbone
if backbone == 'efficientnet_b0':
self.backbone = models.efficientnet_b0(
weights='IMAGENET1K_V1' if pretrained else None
)
feature_dim = 1280
self.backbone.classifier = nn.Identity()
elif backbone == 'resnet18':
self.backbone = models.resnet18(
weights='IMAGENET1K_V1' if pretrained else None
)
feature_dim = 512
self.backbone.fc = nn.Identity()
else:
raise ValueError(f"Unknown backbone: {backbone}")
# Классификатор
self.classifier = nn.Sequential(
nn.Dropout(dropout),
nn.Linear(feature_dim, 256),
nn.ReLU(),
nn.Dropout(dropout),
nn.Linear(256, num_classes)
)
self.feature_dim = feature_dim
def forward(self, x: torch.Tensor) -> torch.Tensor:
features = self.backbone(x)
return self.classifier(features)
def predict(self, x: torch.Tensor) -> Dict:
"""
Предсказание региона.
Args:
x: Input tensor [B, 3, H, W]
Returns:
Dict с keys: predicted_region, probabilities, confidence
"""
self.eval()
with torch.no_grad():
logits = self.forward(x)
probs = torch.softmax(logits, dim=1)
conf, pred = probs.max(dim=1)
results = []
for i in range(x.shape[0]):
region_idx = pred[i].item()
results.append({
'predicted_region': self.REGIONS[region_idx],
'region_index': region_idx,
'confidence': conf[i].item(),
'probabilities': {
self.REGIONS[j]: probs[i, j].item()
for j in range(self.num_classes)
}
})
return results[0] if len(results) == 1 else results
def get_info(self) -> Dict:
return {
'name': 'RegionDetector',
'backbone': self.backbone_name,
'num_classes': self.num_classes,
'regions': self.REGIONS,
'feature_dim': self.feature_dim
}
class RegionDetectorWithFeatures(RegionDetector):
"""Region detector с извлечением признаков"""
def extract_features(self, x: torch.Tensor) -> torch.Tensor:
return self.backbone(x)
def create_region_detector(
backbone: str = 'efficientnet_b0',
num_classes: int = 4,
pretrained: bool = True,
device: str = 'cpu'
) -> RegionDetector:
"""Создание модели"""
model = RegionDetector(
backbone=backbone,
num_classes=num_classes,
pretrained=pretrained
)
model.to(device)
return model
# Функция для определения региона из изображения (rule-based fallback)
def determine_region_from_image(image: np.ndarray) -> str:
"""
Определение региона из изображения без модели.
Использует форму яркой области.
Args:
image: numpy array [H, W] или [H, W, 3]
Returns:
str: 'spine', 'hip_left', 'hip_right', 'hip'
"""
if len(image.shape) == 3:
image = np.mean(image, axis=2)
# Нормализация
image = (image - image.min()) / (image.max() - image.min() + 1e-8)
h, w = image.shape
# Порог для яркой области
threshold = np.percentile(image, 95)
binary = image > threshold
if binary.sum() == 0:
return 'unknown'
# Aspect ratio яркой области
from scipy import ndimage
rows = np.any(binary, axis=1)
cols = np.any(binary, axis=0)
if not (rows.any() and cols.any()):
return 'unknown'
try:
rmin, rmax = np.where(rows)[0][[0, -1]]
cmin, cmax = np.where(cols)[0][[0, -1]]
bbox_h = rmax - rmin
bbox_w = cmax - cmin
bbox_aspect = bbox_h / (bbox_w + 1e-6)
# Симметрия
h_mid, w_mid = h // 2, w // 2
left_half = image[:, :w_mid]
right_half = np.fliplr(image[:, w_mid:])
min_w = min(left_half.shape[1], right_half.shape[1])
symmetry = 1 - np.abs(left_half[:, :min_w] - right_half[:, :min_w]).mean() / (image.std() + 1e-6)
# Left/right brightness ratio
left_bright = binary[:, :w//2].sum()
right_bright = binary[:, w//2:].sum()
left_right_ratio = left_bright / (right_bright + 1e-6)
# Классификация
if bbox_aspect < 1.5:
return 'spine'
elif bbox_aspect < 1.8:
return 'spine' if symmetry > 0.35 else 'hip'
else:
if left_right_ratio > 1.3:
return 'hip_right'
elif left_right_ratio < 0.7:
return 'hip_left'
else:
return 'hip'
except:
return 'unknown'

View File

@ -0,0 +1,10 @@
"""
Segmentation models
"""
from .segmentator import DXASegmenter, DXASegmenterSimple, create_segmenter
__all__ = [
'DXASegmenter',
'DXASegmenterSimple',
'create_segmenter'
]

View File

@ -0,0 +1,254 @@
"""
U-Net сегментация для DXA изображений
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
import torchvision.models as models
from typing import Dict, Tuple, Optional, List
class DoubleConv(nn.Module):
"""Double convolution block"""
def __init__(self, in_channels, out_channels):
super().__init__()
self.conv = nn.Sequential(
nn.Conv2d(in_channels, out_channels, 3, padding=1),
nn.BatchNorm2d(out_channels),
nn.ReLU(inplace=True),
nn.Conv2d(out_channels, out_channels, 3, padding=1),
nn.BatchNorm2d(out_channels),
nn.ReLU(inplace=True)
)
def forward(self, x):
return self.conv(x)
class Down(nn.Module):
"""Downsampling block"""
def __init__(self, in_channels, out_channels):
super().__init__()
self.maxpool_conv = nn.Sequential(
nn.MaxPool2d(2),
DoubleConv(in_channels, out_channels)
)
def forward(self, x):
return self.maxpool_conv(x)
class Up(nn.Module):
"""Upsampling block"""
def __init__(self, in_channels, out_channels):
super().__init__()
self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, 2, stride=2)
self.conv = DoubleConv(in_channels, out_channels)
def forward(self, x1, x2):
x1 = self.up(x1)
# Pad if sizes don't match
diffY = x2.size()[2] - x1.size()[2]
diffX = x2.size()[3] - x1.size()[3]
x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,
diffY // 2, diffY - diffY // 2])
x = torch.cat([x2, x1], dim=1)
return self.conv(x)
class OutConv(nn.Module):
"""Output convolution"""
def __init__(self, in_channels, out_channels):
super().__init__()
self.conv = nn.Conv2d(in_channels, out_channels, 1)
def forward(self, x):
return self.conv(x)
class DXASegmenter(nn.Module):
"""
U-Net сегментатор для DXA изображений.
Может работать для spine (позвонки) и hip (бедро).
"""
def __init__(self,
in_channels: int = 3,
out_classes: int = 2,
base_features: int = 64,
backbone: str = 'resnet34'):
super().__init__()
self.in_channels = in_channels
self.out_classes = out_classes
self.base_features = base_features
# Encoder (используем pretrained backbone)
if backbone == 'resnet34':
resnet = models.resnet34(weights='IMAGENET1K_V1')
self.encoder1 = nn.Sequential(resnet.conv1, resnet.bn1, resnet.relu)
self.encoder2 = nn.Sequential(resnet.maxpool, resnet.layer1)
self.encoder3 = resnet.layer2 # 256
self.encoder4 = resnet.layer3 # 512
self.encoder5 = resnet.layer4 # 1024
encoder_out = 1024
elif backbone == 'resnet18':
resnet = models.resnet18(weights='IMAGENET1K_V1')
self.encoder1 = nn.Sequential(resnet.conv1, resnet.bn1, resnet.relu)
self.encoder2 = nn.Sequential(resnet.maxpool, resnet.layer1)
self.encoder3 = resnet.layer2 # 128
self.encoder4 = resnet.layer3 # 256
self.encoder5 = resnet.layer4 # 512
encoder_out = 512
else:
raise ValueError(f"Unknown backbone: {backbone}")
# Decoder
self.up5 = Up(encoder_out, 512)
self.up4 = Up(512, 256)
self.up3 = Up(256, 128)
self.up2 = Up(128, 64)
self.out = OutConv(64, out_classes)
def forward(self, x: torch.Tensor) -> torch.Tensor:
"""
Forward pass
Args:
x: [B, 3, H, W] Input image
Returns:
[B, out_classes, H, W] Segmentation logits
"""
# Encoder
x1 = self.encoder1(x) # 64
x2 = self.encoder2(x1) # 64
x3 = self.encoder3(x2) # 128/256
x4 = self.encoder4(x3) # 256/512
x5 = self.encoder5(x4) # 512/1024
# Decoder
d5 = self.up5(x5, x4)
d4 = self.up4(d5, x3)
d3 = self.up3(d4, x2)
d2 = self.up2(d3, x1)
return self.out(d2)
def predict(self, x: torch.Tensor) -> Dict:
"""
Предсказание сегментации.
Returns:
Dict с mask, probabilities
"""
self.eval()
with torch.no_grad():
logits = self.forward(x)
probs = torch.softmax(logits, dim=1)
mask = probs.argmax(dim=1)
return {
'mask': mask, # [B, H, W]
'probabilities': probs, # [B, classes, H, W]
'logits': logits
}
def predict_mask(self, x: torch.Tensor) -> torch.Tensor:
"""Получить только маску"""
result = self.predict(x)
return result['mask']
class DXASegmenterSimple(nn.Module):
"""
Упрощённая U-Net архитектура для сегментации.
"""
def __init__(self,
in_channels: int = 3,
out_classes: int = 2,
features: int = 64):
super().__init__()
self.in_channels = in_channels
self.out_classes = out_classes
# Encoder
self.inc = DoubleConv(in_channels, features)
self.down1 = Down(features, features * 2)
self.down2 = Down(features * 2, features * 4)
self.down3 = Down(features * 4, features * 8)
# Bottleneck
self.bottleneck = DoubleConv(features * 8, features * 16)
# Decoder
self.up3 = Up(features * 16, features * 8)
self.up2 = Up(features * 8, features * 4)
self.up1 = Up(features * 4, features * 2)
self.up0 = Up(features * 2, features)
self.out = OutConv(features, out_classes)
def forward(self, x):
x1 = self.inc(x)
x2 = self.down1(x1)
x3 = self.down2(x2)
x4 = self.down3(x3)
xb = self.bottleneck(x4)
d3 = self.up3(xb, x4)
d2 = self.up2(d3, x3)
d1 = self.up1(d2, x2)
d0 = self.up0(d1, x1)
return self.out(d0)
def predict(self, x: torch.Tensor) -> Dict:
self.eval()
with torch.no_grad():
logits = self.forward(x)
probs = torch.softmax(logits, dim=1)
mask = probs.argmax(dim=1)
return {
'mask': mask,
'probabilities': probs
}
def create_segmenter(
architecture: str = 'unet_resnet34',
in_channels: int = 3,
out_classes: int = 2,
device: str = 'cpu'
) -> nn.Module:
"""
Создание сегментатора.
Args:
architecture: 'unet_resnet34', 'unet_resnet18', 'unet_simple'
in_channels: количество входных каналов
out_classes: количество классов сегментации
device: устройство
"""
if architecture == 'unet_resnet34':
model = DXASegmenter(in_channels, out_classes, backbone='resnet34')
elif architecture == 'unet_resnet18':
model = DXASegmenter(in_channels, out_classes, backbone='resnet18')
elif architecture == 'unet_simple':
model = DXASegmenterSimple(in_channels, out_classes)
else:
raise ValueError(f"Unknown architecture: {architecture}")
model.to(device)
return model

View File

@ -0,0 +1,10 @@
"""
Visualization models
"""
from .gradcam import GradCAMExtractor, SimpleGradCAM, visualize_attention
__all__ = [
'GradCAMExtractor',
'SimpleGradCAM',
'visualize_attention'
]

View File

@ -0,0 +1,249 @@
"""
Grad-CAM для визуализации attention модели
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from typing import Dict, Optional, Tuple
import cv2
class GradCAMExtractor:
"""
Извлечение Grad-CAM heatmap для визуализации внимания модели.
"""
def __init__(self,
model: nn.Module,
target_layer: Optional[nn.Module] = None,
target_class: Optional[int] = None):
"""
Args:
model: Модель для визуализации
target_layer: Слой для извлечения Grad-CAM
target_class: Класс для визуализации (None = pred class)
"""
self.model = model
self.target_layer = target_layer
self.target_class = target_class
self.gradients = None
self.activations = None
# Регистрация hooks
self._register_hooks()
def _register_hooks(self):
"""Регистрация forward и backward hooks"""
def forward_hook(module, input, output):
self.activations = output.detach()
def backward_hook(module, grad_input, grad_output):
self.gradients = grad_output[0].detach()
if self.target_layer is not None:
self.target_layer.register_forward_hook(forward_hook)
self.target_layer.register_full_backward_hook(backward_hook)
def set_target_layer(self, layer: nn.Module):
"""Установка целевого слоя"""
self.target_layer = layer
self._register_hooks()
def generate_cam(self,
input_tensor: torch.Tensor,
target_class: Optional[int] = None) -> np.ndarray:
"""
Генерация CAM heatmap.
Args:
input_tensor: [1, C, H, W] Input image
target_class: Класс для визуализации
Returns:
Heatmap [H, W] normalized 0-1
"""
self.model.eval()
# Forward
output = self.model(input_tensor)
# Выбор класса
if target_class is None:
target_class = output.argmax(dim=1).item()
# Backward
self.model.zero_grad()
class_loss = output[0, target_class]
class_loss.backward()
# Вычисление CAM
if self.gradients is None or self.activations is None:
raise RuntimeError("Gradients or activations not captured")
# Global average pooling градиентов
pooled_gradients = torch.mean(self.gradients, dim=[0, 2, 3])
# Умножение активаций на градиенты
for i in range(self.activations.shape[1]):
self.activations[:, i, :, :] *= pooled_gradients[i]
# Average и ReLU
heatmap = torch.mean(self.activations, dim=1).squeeze()
heatmap = F.relu(heatmap)
# Normalization
if heatmap.max() > 0:
heatmap = heatmap / heatmap.max()
return heatmap.cpu().numpy()
def generate_cam_on_image(self,
input_tensor: torch.Tensor,
original_image: np.ndarray,
target_class: Optional[int] = None,
colormap: int = cv2.COLORMAP_JET) -> np.ndarray:
"""
Генерация CAM и наложение на оригинальное изображение.
Args:
input_tensor: [1, C, H, W] Input tensor
original_image: [H, W, 3] RGB image (0-255)
target_class: Класс для визуализации
colormap: OpenCV colormap
Returns:
Overlay image [H, W, 3]
"""
# Get CAM
heatmap = self.generate_cam(input_tensor, target_class)
# Resize to original image
heatmap = cv2.resize(heatmap, (original_image.shape[1], original_image.shape[0]))
# Apply colormap
heatmap_colored = cv2.applyColorMap((heatmap * 255).astype(np.uint8), colormap)
heatmap_colored = cv2.cvtColor(heatmap_colored, cv2.COLOR_BGR2RGB)
# Overlay
overlay = cv2.addWeighted(original_image, 0.6, heatmap_colored, 0.4, 0)
return overlay
class SimpleGradCAM:
"""
Упрощённый Grad-CAM без hook - работает с любым классификатором.
"""
@staticmethod
def generate(model: nn.Module,
input_tensor: torch.Tensor,
target_layer_name: str = 'encoder5') -> np.ndarray:
"""
Генерация heatmap.
Args:
model: Модель (должна иметь get_info() или быть классификатором)
input_tensor: [1, C, H, W]
target_layer_name: имя слоя для CAM
Returns:
Heatmap [H, W]
"""
model.eval()
# Для ResNet - последний слой
if hasattr(model, 'backbone'):
# Это наша модель
target_layer = model.backbone.layer4
else:
# Ищем последний conv слой
target_layer = None
for name, module in model.named_modules():
if isinstance(module, nn.Conv2d):
target_layer = module
if target_layer is None:
raise ValueError("Cannot find target layer")
extractor = GradCAMExtractor(model, target_layer)
# Forward + backward
output = model(input_tensor)
target_class = output.argmax(dim=1).item()
model.zero_grad()
output[0, target_class].backward()
# CAM
gradients = extractor.gradients
activations = extractor.activations
if gradients is None or activations is None:
raise RuntimeError("Failed to capture gradients")
pooled_gradients = torch.mean(gradients, dim=[0, 2, 3])
for i in range(activations.shape[1]):
activations[:, i, :, :] *= pooled_gradients[i]
heatmap = torch.mean(activations, dim=1).squeeze()
heatmap = F.relu(heatmap)
if heatmap.max() > 0:
heatmap = heatmap / heatmap.max()
return heatmap.cpu().numpy()
def visualize_attention(model: nn.Module,
image: torch.Tensor,
original_rgb: np.ndarray,
target_class: Optional[int] = None) -> Tuple[np.ndarray, np.ndarray]:
"""
Утилита для визуализации attention.
Args:
model: Классификатор
image: [1, 3, H, W] Tensor
original_rgb: [H, W, 3] RGB numpy
target_class: Опционально класс для визуализации
Returns:
(heatmap, overlay) - heatmap и наложенное изображение
"""
# Try to use Grad-CAM
try:
# Find target layer
target_layer = None
for name, module in model.named_modules():
if hasattr(module, 'out_channels') and isinstance(module, nn.Conv2d):
target_layer = module
if target_layer is None:
# Use last conv layer of backbone
if hasattr(model, 'backbone'):
target_layer = model.backbone[-1]
if target_layer is not None:
extractor = GradCAMExtractor(model, target_layer)
heatmap = extractor.generate_cam(image, target_class)
# Resize
heatmap = cv2.resize(heatmap, (original_rgb.shape[1], original_rgb.shape[0]))
# Overlay
heatmap_colored = cv2.applyColorMap(
(heatmap * 255).astype(np.uint8),
cv2.COLORMAP_JET
)
heatmap_colored = cv2.cvtColor(heatmap_colored, cv2.COLOR_BGR2RGB)
overlay = cv2.addWeighted(original_rgb, 0.6, heatmap_colored, 0.4, 0)
return heatmap, overlay
except Exception as e:
print(f"Grad-CAM failed: {e}")
# Fallback: return original image
return np.zeros((original_rgb.shape[0], original_rgb.shape[1])), original_rgb

16
src/pipeline/__init__.py Normal file
View File

@ -0,0 +1,16 @@
"""
Pipeline components
"""
from .pipeline import (
QualityAssessmentPipeline,
PipelineConfig,
QualityResult,
create_pipeline
)
__all__ = [
'QualityAssessmentPipeline',
'PipelineConfig',
'QualityResult',
'create_pipeline'
]

289
src/pipeline/pipeline.py Normal file
View File

@ -0,0 +1,289 @@
"""
Основной пайплайн оценки качества DXA
"""
import torch
import numpy as np
from PIL import Image
from typing import Dict, List, Optional, Any
from dataclasses import dataclass, field
import pydicom
from src.models import (
RegionDetector,
QualityClassifier,
ViolationTypeClassifier,
DXASegmenter,
create_segmenter,
determine_region_from_image
)
from src.utils.utils import get_device
@dataclass
class PipelineConfig:
"""Конфигурация пайплайна"""
# Paths to models
region_detector_path: Optional[str] = None
quality_classifier_path: Optional[str] = None
violation_classifier_path: Optional[str] = None
segmentator_path: Optional[str] = None
# Model configs
backbone: str = 'resnet18'
input_size: int = 224
device: str = 'auto'
# Режимы
use_segmentation: bool = True
use_violation_classification: bool = True
use_attention: bool = False
# Threshold
quality_threshold: float = 0.5
confidence_threshold: float = 0.7
@dataclass
class QualityResult:
"""Результат оценки качества"""
# Идентификация
study_uid: str = ""
image_uid: str = ""
filename: str = ""
# Основные результаты
anatomical_region: str = "unknown"
quality_class: int = 0 # 0 = OK, 1 = Violation
quality_label: str = "OK"
# Детали
violation_type: str = "correct"
violation_description: str = ""
reason: str = ""
# Confidence
confidence: float = 0.0
confidence_per_class: Dict[str, float] = field(default_factory=dict)
# View quality
view_quality: str = "unknown"
# Метрики
metrics: Dict[str, Any] = field(default_factory=dict)
# Визуализация
mask: Optional[np.ndarray] = None
heatmap: Optional[np.ndarray] = None
# Статус
processing_status: str = "Success"
error: Optional[str] = None
class QualityAssessmentPipeline:
"""
Основной пайплайн для оценки качества DXA исследований.
"""
def __init__(self, config: PipelineConfig):
self.config = config
self.device_str = config.device if config.device != 'auto' else get_device()
self.device = torch.device(self.device_str)
# Модели (загружаются лениво)
self._region_detector: Optional[RegionDetector] = None
self._quality_classifier: Optional[QualityClassifier] = None
self._violation_classifier: Optional[ViolationTypeClassifier] = None
self._segmentator: Optional[DXASegmenter] = None
@property
def region_detector(self) -> RegionDetector:
if self._region_detector is None:
self._region_detector = RegionDetector(
backbone=self.config.backbone,
pretrained=True
).to(self.device)
self._region_detector.eval()
return self._region_detector
@property
def quality_classifier(self) -> QualityClassifier:
if self._quality_classifier is None:
self._quality_classifier = QualityClassifier(
backbone=self.config.backbone,
pretrained=True,
use_attention=self.config.use_attention
).to(self.device)
self._quality_classifier.eval()
return self._quality_classifier
@property
def violation_classifier(self) -> ViolationTypeClassifier:
if self._violation_classifier is None:
self._violation_classifier = ViolationTypeClassifier(
backbone=self.config.backbone,
pretrained=True
).to(self.device)
self._violation_classifier.eval()
return self._violation_classifier
@property
def segmentator(self) -> DXASegmenter:
if self._segmentator is None:
self._segmentator = create_segmenter(
architecture='unet_resnet18',
in_channels=3,
out_classes=2,
device=self.device_str
)
self._segmentator.eval()
return self._segmentator
def _preprocess_image(self, image: np.ndarray) -> torch.Tensor:
"""Предобработка изображения для модели"""
# Нормализация
if image.max() > 1:
image = image.astype(np.float32) / 255.0
# Grayscale -> RGB
if len(image.shape) == 2:
image = np.stack([image] * 3, axis=2)
elif image.shape[2] == 1:
image = np.concatenate([image] * 3, axis=2)
# Resize
if isinstance(self.config.input_size, int):
h, w = image.shape[:2]
if h != self.config.input_size or w != self.config.input_size:
img_pil = Image.fromarray((image * 255).astype(np.uint8))
img_pil = img_pil.resize(
(self.config.input_size, self.config.input_size),
Image.BILINEAR
)
image = np.array(img_pil).astype(np.float32) / 255.0
# CHW
image = image.transpose(2, 0, 1)
return torch.from_numpy(image).unsqueeze(0).to(self.device)
def analyze(self, dicom_path: str) -> QualityResult:
"""
Полный анализ DICOM файла.
Args:
dicom_path: Путь к DICOM файлу
Returns:
QualityResult с результатами оценки
"""
result = QualityResult()
try:
# Загрузка DICOM
ds = pydicom.dcmread(dicom_path)
image = ds.pixel_array.astype(np.float32)
# UID
result.study_uid = getattr(ds, 'StudyInstanceUID', '')
result.image_uid = getattr(ds, 'SOPInstanceUID', '')
result.filename = dicom_path
# Нормализация
image = (image - image.min()) / (image.max() - image.min() + 1e-8)
# Определение региона (сначала rule-based, потом модель)
region = determine_region_from_image(image)
result.anatomical_region = region
# Предобработка для моделей
input_tensor = self._preprocess_image(image)
# 1. Классификация качества
quality_result = self.quality_classifier.predict(input_tensor)
result.quality_class = quality_result['predicted_class']
result.quality_label = quality_result['label']
result.confidence = quality_result['confidence']
result.confidence_per_class = quality_result['probabilities']
# 2. Определение типа нарушения (если есть)
if result.quality_class == 1 and self.config.use_violation_classification:
violation_result = self.violation_classifier.predict(
input_tensor,
region='spine' if region == 'spine' else 'hip'
)
result.violation_type = violation_result['predicted_type']
result.violation_description = violation_result['description']
result.reason = violation_result['description']
# 3. Сегментация (опционально)
if self.config.use_segmentation:
try:
seg_result = self.segmentator.predict(input_tensor)
result.mask = seg_result['mask'].cpu().numpy()[0]
except Exception as e:
result.metrics['segmentation_error'] = str(e)
# Метрики
result.metrics = {
'device': self.device_str,
'backbone': self.config.backbone,
'input_size': self.config.input_size,
'use_segmentation': self.config.use_segmentation,
'use_violation': self.config.use_violation_classification
}
except Exception as e:
result.processing_status = "Failure"
result.error = str(e)
return result
def analyze_batch(self, dicom_paths: List[str]) -> List[QualityResult]:
"""Анализ нескольких файлов"""
return [self.analyze(path) for path in dicom_paths]
def to_dict(self, result: QualityResult) -> Dict:
"""Конвертация результата в словарь для JSON"""
d = {
'study_uid': result.study_uid,
'image_uid': result.image_uid,
'filename': result.filename,
'anatomical_region': result.anatomical_region,
'quality_class': result.quality_class,
'quality_label': result.quality_label,
'violation_type': result.violation_type,
'violation_description': result.violation_description,
'reason': result.reason,
'confidence': result.confidence,
'confidence_per_class': result.confidence_per_class,
'view_quality': result.view_quality,
'metrics': result.metrics,
'processing_status': result.processing_status,
'error': result.error
}
# Convert numpy types
def convert(obj):
if isinstance(obj, np.ndarray):
return obj.tolist()
elif isinstance(obj, (np.integer, np.floating)):
return obj.item()
elif isinstance(obj, dict):
return {k: convert(v) for k, v in obj.items()}
elif isinstance(obj, list):
return [convert(i) for i in obj]
return obj
return convert(d)
def create_pipeline(config: Optional[PipelineConfig] = None,
device: Optional[str] = None) -> QualityAssessmentPipeline:
"""Создание пайплайна"""
if config is None:
config = PipelineConfig()
if device is not None:
config.device = device
return QualityAssessmentPipeline(config)

281
src/training/train_model.py Normal file
View File

@ -0,0 +1,281 @@
"""
Training script for DXA Quality Models
"""
import os
import sys
import argparse
from pathlib import Path
import numpy as np
import pandas as pd
import torch
import torch.nn as nn
from torch.utils.data import Dataset, DataLoader, Subset
from PIL import Image
import pydicom
from tqdm import tqdm
import random
# Add project root to path
project_root = Path(__file__).parent.parent.parent
sys.path.insert(0, str(project_root))
os.chdir(project_root)
class DXATrainingDataset(Dataset):
"""Dataset for training DXA models"""
def __init__(self, data_root, annotation_path, region=None, input_size=224):
self.data_root = Path(data_root)
self.annotation_path = annotation_path
self.region = region
self.input_size = input_size
self.annotation = self._load_annotation()
self.samples = self._build_samples()
print(f"Loaded {len(self.samples)} samples for region={region}")
def _load_annotation(self):
df = pd.read_excel(self.annotation_path, header=None)
data = df.iloc[2:].copy()
data.columns = range(len(df.columns))
data = data.rename(columns={
0: 'id',
1: 'study_uid',
9: 'spine_total',
10: 'hip_right_total',
11: 'hip_left_total'
})
data = data.dropna(subset=['study_uid'])
return data
def _build_samples(self):
samples = []
for _, row in self.annotation.iterrows():
study_uid = str(row['study_uid']).strip()
# Find study folder - it's inside Исследования subfolder
study_path = self.data_root / 'Исследования' / study_uid
if not study_path.exists():
continue
# Find DICOM files
dcm_files = sorted(study_path.rglob('*.dcm'))
for dcm_file in dcm_files:
fname = dcm_file.name.lower()
# Determine region from filename
if 'spine' in fname:
sample_region = 'spine'
target_col = 'spine_total'
elif 'l_hip' in fname or 'left' in fname:
sample_region = 'hip_left'
target_col = 'hip_left_total'
elif 'r_hip' in fname or 'right' in fname:
sample_region = 'hip_right'
target_col = 'hip_right_total'
else:
continue
# Filter by region if specified
if self.region and sample_region != self.region:
continue
# Get target from the same row
if target_col in row and pd.notna(row[target_col]):
target = int(row[target_col])
else:
# Use spine total as fallback
if pd.notna(row.get('spine_total', None)):
target = int(row['spine_total'])
else:
continue
samples.append({
'path': str(dcm_file),
'region': sample_region,
'target': target
})
return samples
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
sample = self.samples[idx]
# Load DICOM
ds = pydicom.dcmread(sample['path'])
img = ds.pixel_array.astype(np.float32)
# Normalize
img = (img - img.min()) / (img.max() - img.min() + 1e-8)
# Convert to RGB
img = np.stack([img] * 3, axis=2)
# Resize
img_pil = Image.fromarray((img * 255).astype(np.uint8))
img_pil = img_pil.resize((self.input_size, self.input_size), Image.BILINEAR)
img = np.array(img_pil).astype(np.float32) / 255.0
# CHW
img = img.transpose(2, 0, 1)
img = torch.from_numpy(img).float()
return img, torch.tensor(sample['target'], dtype=torch.long)
def train_model(model, train_loader, val_loader, epochs, lr, device, save_path):
"""Train model"""
model = model.to(device)
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=2)
best_val_acc = 0
for epoch in range(epochs):
# Train
model.train()
train_loss = 0
train_correct = 0
train_total = 0
pbar = tqdm(train_loader, desc=f"Epoch {epoch+1}/{epochs}")
for images, labels in pbar:
images = images.to(device)
labels = labels.to(device)
optimizer.zero_grad()
outputs = model(images)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
train_loss += loss.item() * images.size(0)
_, predicted = outputs.max(1)
train_correct += predicted.eq(labels).sum().item()
train_total += labels.size(0)
pbar.set_postfix({'loss': train_loss/train_total, 'acc': train_correct/train_total})
train_acc = train_correct / train_total
# Validate
model.eval()
val_loss = 0
val_correct = 0
val_total = 0
with torch.no_grad():
for images, labels in val_loader:
images = images.to(device)
labels = labels.to(device)
outputs = model(images)
loss = criterion(outputs, labels)
val_loss += loss.item() * images.size(0)
_, predicted = outputs.max(1)
val_correct += predicted.eq(labels).sum().item()
val_total += labels.size(0)
val_acc = val_correct / val_total
scheduler.step(val_loss)
print(f"Epoch {epoch+1}: Train Acc={train_acc:.3f}, Val Acc={val_acc:.3f}, Val Loss={val_loss/val_total:.4f}")
if val_acc > best_val_acc:
best_val_acc = val_acc
torch.save(model.state_dict(), save_path)
print(f" -> Saved best model to {save_path}")
return best_val_acc
def main():
parser = argparse.ArgumentParser()
parser.add_argument('--data-root', default='dataset_hack/НД_для_обучения')
parser.add_argument('--annotation', default='dataset_hack/НД_для_обучения/разметка.xlsx')
parser.add_argument('--model-type', default='quality', choices=['quality', 'region', 'violation'])
parser.add_argument('--epochs', type=int, default=15)
parser.add_argument('--batch-size', type=int, default=8)
parser.add_argument('--lr', type=float, default=1e-4)
parser.add_argument('--input-size', type=int, default=224)
parser.add_argument('--output', default='models/')
parser.add_argument('--region', default=None, help='spine, hip_left, hip_right, or None for all')
args = parser.parse_args()
# Device
device = torch.device('cuda' if torch.cuda.is_available() else 'mps' if torch.backends.mps.is_available() else 'cpu')
print(f"Using device: {device}")
# Create output dir
os.makedirs(args.output, exist_ok=True)
# Load dataset
dataset = DXATrainingDataset(
data_root=args.data_root,
annotation_path=args.annotation,
region=args.region,
input_size=args.input_size
)
if len(dataset) == 0:
print("ERROR: No samples found!")
return
# Split
n = len(dataset)
n_train = int(n * 0.8)
n_val = n - n_train
# Random split
indices = list(range(n))
random.shuffle(indices)
train_indices = indices[:n_train]
val_indices = indices[n_train:]
train_dataset = Subset(dataset, train_indices)
val_dataset = Subset(dataset, val_indices)
train_loader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True, num_workers=0)
val_loader = DataLoader(val_dataset, batch_size=args.batch_size, shuffle=False, num_workers=0)
print(f"Train: {len(train_dataset)}, Val: {len(val_dataset)}")
# Class distribution
targets = [dataset.samples[i]['target'] for i in train_indices]
n_pos = sum(targets)
n_neg = len(targets) - n_pos
print(f"Class distribution - Positive: {n_pos}, Negative: {n_neg}")
# Create model
if args.model_type == 'quality':
from src.models.classification import QualityClassifier
model = QualityClassifier(backbone='resnet18', pretrained=True)
save_path = f"{args.output}/quality_classifier.pth"
elif args.model_type == 'region':
from src.models import RegionDetector
model = RegionDetector(backbone='resnet18', pretrained=True)
save_path = f"{args.output}/region_detector.pth"
else:
from src.models.classification import ViolationTypeClassifier
model = ViolationTypeClassifier(backbone='resnet18', pretrained=True)
save_path = f"{args.output}/violation_classifier.pth"
# Train
print(f"Training {args.model_type} model...")
best_acc = train_model(model, train_loader, val_loader, args.epochs, args.lr, device, save_path)
print(f"Best validation accuracy: {best_acc:.3f}")
print(f"Model saved to: {save_path}")
if __name__ == '__main__':
main()