""" Основной пайплайн оценки качества DXA """ import torch import numpy as np from PIL import Image from typing import Dict, List, Optional, Any from dataclasses import dataclass, field import pydicom from src.models import ( RegionDetector, QualityClassifier, ViolationTypeClassifier, DXASegmenter, create_segmenter, determine_region_from_image ) from src.utils.utils import get_device @dataclass class PipelineConfig: """Конфигурация пайплайна""" # Paths to models region_detector_path: Optional[str] = None quality_classifier_path: Optional[str] = None violation_classifier_path: Optional[str] = None segmentator_path: Optional[str] = None # Model configs backbone: str = 'resnet18' input_size: int = 224 device: str = 'auto' # Режимы use_segmentation: bool = True use_violation_classification: bool = True use_attention: bool = False # Threshold quality_threshold: float = 0.5 confidence_threshold: float = 0.7 @dataclass class QualityResult: """Результат оценки качества""" # Идентификация study_uid: str = "" image_uid: str = "" filename: str = "" # Основные результаты anatomical_region: str = "unknown" quality_class: int = 0 # 0 = OK, 1 = Violation quality_label: str = "OK" # Детали violation_type: str = "correct" violation_description: str = "" reason: str = "" # Confidence confidence: float = 0.0 confidence_per_class: Dict[str, float] = field(default_factory=dict) # View quality view_quality: str = "unknown" # Метрики metrics: Dict[str, Any] = field(default_factory=dict) # Визуализация mask: Optional[np.ndarray] = None heatmap: Optional[np.ndarray] = None # Статус processing_status: str = "Success" error: Optional[str] = None class QualityAssessmentPipeline: """ Основной пайплайн для оценки качества DXA исследований. """ def __init__(self, config: PipelineConfig): self.config = config self.device_str = config.device if config.device != 'auto' else get_device() self.device = torch.device(self.device_str) # Модели (загружаются лениво) self._region_detector: Optional[RegionDetector] = None self._quality_classifier: Optional[QualityClassifier] = None self._violation_classifier: Optional[ViolationTypeClassifier] = None self._segmentator: Optional[DXASegmenter] = None @property def region_detector(self) -> RegionDetector: if self._region_detector is None: self._region_detector = RegionDetector( backbone=self.config.backbone, pretrained=True ).to(self.device) self._region_detector.eval() return self._region_detector @property def quality_classifier(self) -> QualityClassifier: if self._quality_classifier is None: self._quality_classifier = QualityClassifier( backbone=self.config.backbone, pretrained=True, use_attention=self.config.use_attention ).to(self.device) self._quality_classifier.eval() return self._quality_classifier @property def violation_classifier(self) -> ViolationTypeClassifier: if self._violation_classifier is None: self._violation_classifier = ViolationTypeClassifier( backbone=self.config.backbone, pretrained=True ).to(self.device) self._violation_classifier.eval() return self._violation_classifier @property def segmentator(self) -> DXASegmenter: if self._segmentator is None: self._segmentator = create_segmenter( architecture='unet_resnet18', in_channels=3, out_classes=2, device=self.device_str ) self._segmentator.eval() return self._segmentator def _preprocess_image(self, image: np.ndarray) -> torch.Tensor: """Предобработка изображения для модели""" # Нормализация if image.max() > 1: image = image.astype(np.float32) / 255.0 # Grayscale -> RGB if len(image.shape) == 2: image = np.stack([image] * 3, axis=2) elif image.shape[2] == 1: image = np.concatenate([image] * 3, axis=2) # Resize if isinstance(self.config.input_size, int): h, w = image.shape[:2] if h != self.config.input_size or w != self.config.input_size: img_pil = Image.fromarray((image * 255).astype(np.uint8)) img_pil = img_pil.resize( (self.config.input_size, self.config.input_size), Image.BILINEAR ) image = np.array(img_pil).astype(np.float32) / 255.0 # CHW image = image.transpose(2, 0, 1) return torch.from_numpy(image).unsqueeze(0).to(self.device) def analyze(self, dicom_path: str) -> QualityResult: """ Полный анализ DICOM файла. Args: dicom_path: Путь к DICOM файлу Returns: QualityResult с результатами оценки """ result = QualityResult() try: # Загрузка DICOM ds = pydicom.dcmread(dicom_path) image = ds.pixel_array.astype(np.float32) # UID result.study_uid = getattr(ds, 'StudyInstanceUID', '') result.image_uid = getattr(ds, 'SOPInstanceUID', '') result.filename = dicom_path # Нормализация image = (image - image.min()) / (image.max() - image.min() + 1e-8) # Определение региона (сначала rule-based, потом модель) region = determine_region_from_image(image) result.anatomical_region = region # Предобработка для моделей input_tensor = self._preprocess_image(image) # 1. Классификация качества quality_result = self.quality_classifier.predict(input_tensor) result.quality_class = quality_result['predicted_class'] result.quality_label = quality_result['label'] result.confidence = quality_result['confidence'] result.confidence_per_class = quality_result['probabilities'] # 2. Определение типа нарушения (если есть) if result.quality_class == 1 and self.config.use_violation_classification: violation_result = self.violation_classifier.predict( input_tensor, region='spine' if region == 'spine' else 'hip' ) result.violation_type = violation_result['predicted_type'] result.violation_description = violation_result['description'] result.reason = violation_result['description'] # 3. Сегментация (опционально) if self.config.use_segmentation: try: seg_result = self.segmentator.predict(input_tensor) result.mask = seg_result['mask'].cpu().numpy()[0] except Exception as e: result.metrics['segmentation_error'] = str(e) # Метрики result.metrics = { 'device': self.device_str, 'backbone': self.config.backbone, 'input_size': self.config.input_size, 'use_segmentation': self.config.use_segmentation, 'use_violation': self.config.use_violation_classification } except Exception as e: result.processing_status = "Failure" result.error = str(e) return result def analyze_batch(self, dicom_paths: List[str]) -> List[QualityResult]: """Анализ нескольких файлов""" return [self.analyze(path) for path in dicom_paths] def to_dict(self, result: QualityResult) -> Dict: """Конвертация результата в словарь для JSON""" d = { 'study_uid': result.study_uid, 'image_uid': result.image_uid, 'filename': result.filename, 'anatomical_region': result.anatomical_region, 'quality_class': result.quality_class, 'quality_label': result.quality_label, 'violation_type': result.violation_type, 'violation_description': result.violation_description, 'reason': result.reason, 'confidence': result.confidence, 'confidence_per_class': result.confidence_per_class, 'view_quality': result.view_quality, 'metrics': result.metrics, 'processing_status': result.processing_status, 'error': result.error } # Convert numpy types def convert(obj): if isinstance(obj, np.ndarray): return obj.tolist() elif isinstance(obj, (np.integer, np.floating)): return obj.item() elif isinstance(obj, dict): return {k: convert(v) for k, v in obj.items()} elif isinstance(obj, list): return [convert(i) for i in obj] return obj return convert(d) def create_pipeline(config: Optional[PipelineConfig] = None, device: Optional[str] = None) -> QualityAssessmentPipeline: """Создание пайплайна""" if config is None: config = PipelineConfig() if device is not None: config.device = device return QualityAssessmentPipeline(config)