develop - hack_2026
This commit is contained in:
parent
c40cde7ec0
commit
3121ecb8a3
22
QWEN.md
22
QWEN.md
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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 |
|
||||
|
||||
---
|
||||
|
||||
*Проект разработан в рамках хакатона по медицинскому ИИ*
|
||||
|
|
@ -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
|
||||
|
|
@ -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'
|
||||
]
|
||||
|
|
@ -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)
|
||||
|
|
@ -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'
|
||||
]
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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'
|
||||
|
|
@ -0,0 +1,10 @@
|
|||
"""
|
||||
Segmentation models
|
||||
"""
|
||||
from .segmentator import DXASegmenter, DXASegmenterSimple, create_segmenter
|
||||
|
||||
__all__ = [
|
||||
'DXASegmenter',
|
||||
'DXASegmenterSimple',
|
||||
'create_segmenter'
|
||||
]
|
||||
|
|
@ -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
|
||||
|
|
@ -0,0 +1,10 @@
|
|||
"""
|
||||
Visualization models
|
||||
"""
|
||||
from .gradcam import GradCAMExtractor, SimpleGradCAM, visualize_attention
|
||||
|
||||
__all__ = [
|
||||
'GradCAMExtractor',
|
||||
'SimpleGradCAM',
|
||||
'visualize_attention'
|
||||
]
|
||||
|
|
@ -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
|
||||
|
|
@ -0,0 +1,16 @@
|
|||
"""
|
||||
Pipeline components
|
||||
"""
|
||||
from .pipeline import (
|
||||
QualityAssessmentPipeline,
|
||||
PipelineConfig,
|
||||
QualityResult,
|
||||
create_pipeline
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
'QualityAssessmentPipeline',
|
||||
'PipelineConfig',
|
||||
'QualityResult',
|
||||
'create_pipeline'
|
||||
]
|
||||
|
|
@ -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)
|
||||
|
|
@ -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()
|
||||
Loading…
Reference in New Issue