bone_2026/src/pipeline/pipeline.py

290 lines
10 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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