548 lines
20 KiB
Markdown
548 lines
20 KiB
Markdown
# Архитектура мультимодельного пайплайна
|
||
|
||
## Обзор
|
||
|
||
Предлагается архитектура из нескольких специализированных моделей, объединённых в единый пайплайн обработки.
|
||
|
||
```
|
||
┌─────────────────────────────────────────────────────────────────────────────┐
|
||
│ 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
|