290 lines
10 KiB
Python
290 lines
10 KiB
Python
"""
|
||
Основной пайплайн оценки качества 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)
|