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