20 KiB
Архитектура мультимодельного пайплайна
Обзор
Предлагается архитектура из нескольких специализированных моделей, объединённых в единый пайплайн обработки.
┌─────────────────────────────────────────────────────────────────────────────┐
│ 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
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)
Обучение:
python src/models/train_region_detector.py --epochs 20
2. Segmentation Model (Сегментация ROI)
Назначение: Сегментация костных структур для дальнейшего анализа
Варианты:
Вариант A: Универсальная модель (рекомендуемый)
| Регион | Выход |
|---|---|
| Spine | Маска позвонков L1-L4 |
| Hip | Маска бедренной кости + шейка + вертелы |
Архитектура: U-Net с ResNet34 encoder
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
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
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:
# 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 (Детектор артефактов)
Назначение: Специализированная детекция артефактов
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 (Модель внимания)
Назначение: Визуализация областей, на которые обращает внимание модель
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
Пайплайн обработки
Основной класс
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
Конфигурация
@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 # Анализ ошибок
Обучение
Скрипт обучения
# 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: Основные модели
- ✅ Region Detector (уже есть в inference.py)
- Quality Classifier (улучшить)
- Violation Type Classifier (новый)
Фаза 2: Сегментация
- Spine Segmentator
- Hip Segmentator
Фаза 3: Улучшения
- Artifact Detector
- Grad-CAM Heatmap
- Multi-task learning