bone_2026/docs/multi_model_architecture.md

548 lines
20 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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