bone_2026/docs/multi_model_architecture.md

20 KiB
Raw Blame History

Архитектура мультимодельного пайплайна

Обзор

Предлагается архитектура из нескольких специализированных моделей, объединённых в единый пайплайн обработки.

┌─────────────────────────────────────────────────────────────────────────────┐
│                        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: Основные модели

  1. ✅ Region Detector (уже есть в inference.py)
  2. Quality Classifier (улучшить)
  3. Violation Type Classifier (новый)

Фаза 2: Сегментация

  1. Spine Segmentator
  2. Hip Segmentator

Фаза 3: Улучшения

  1. Artifact Detector
  2. Grad-CAM Heatmap
  3. Multi-task learning