develop - hack_2026
This commit is contained in:
parent
c40cde7ec0
commit
3121ecb8a3
22
QWEN.md
22
QWEN.md
|
|
@ -214,6 +214,28 @@ The system automatically determines the anatomical region from the DICOM image c
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
## Multi-Model Architecture (Planned)
|
||||||
|
|
||||||
|
See `docs/multi_model_architecture.md` for the planned pipeline:
|
||||||
|
|
||||||
|
```
|
||||||
|
Pipeline:
|
||||||
|
1. Region Detector → 2. Segmentator → 3. Quality Classifier → 4. Violation Type → 5. Aggregator
|
||||||
|
```
|
||||||
|
|
||||||
|
### Planned Models:
|
||||||
|
|
||||||
|
| Model | Purpose | File |
|
||||||
|
|-------|---------|------|
|
||||||
|
| Region Detector | Spine/Hip detection | `src/models/region_detector.py` |
|
||||||
|
| Segmentator | Bone segmentation | `src/models/segmentation/` |
|
||||||
|
| Quality Classifier | OK/Violation binary | `src/models/classification/quality.py` |
|
||||||
|
| Violation Classifier | 7+ violation types | `src/models/classification/violation.py` |
|
||||||
|
| Artifact Detector | Motion, metal detection | `src/models/artifacts/detector.py` |
|
||||||
|
| Grad-CAM | Attention heatmap | `src/models/visualization/gradcam.py` |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
## Docker
|
## Docker
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
|
|
||||||
|
|
@ -0,0 +1,547 @@
|
||||||
|
# Архитектура мультимодельного пайплайна
|
||||||
|
|
||||||
|
## Обзор
|
||||||
|
|
||||||
|
Предлагается архитектура из нескольких специализированных моделей, объединённых в единый пайплайн обработки.
|
||||||
|
|
||||||
|
```
|
||||||
|
┌─────────────────────────────────────────────────────────────────────────────┐
|
||||||
|
│ 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
|
||||||
|
|
@ -0,0 +1,393 @@
|
||||||
|
# DXA Quality Assessment — Документация для презентации
|
||||||
|
|
||||||
|
## Содержание
|
||||||
|
|
||||||
|
1. [Обзор проекта](#1-обзор-проекта)
|
||||||
|
2. [Проблема и актуальность](#2-проблема-и-актуальность)
|
||||||
|
3. [Техническая архитектура](#3-техническая-архитектура)
|
||||||
|
4. [Датасет и разметка](#4-датасет-и-разметка)
|
||||||
|
5. [ML-модель](#5-ml-модель)
|
||||||
|
6. [API и интеграция](#6-api-и-интеграция)
|
||||||
|
7. [Результаты](#7-результаты)
|
||||||
|
8. [Демонстрация](#8-демонстрация)
|
||||||
|
9. [Перспективы развития](#9-перспективы-развития)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 1. Обзор проекта
|
||||||
|
|
||||||
|
### Что это за проект?
|
||||||
|
|
||||||
|
**DXA Quality Assessment** — это AI-сервис для автоматизированной оценки качества денситометрических исследований (DXA — Dual-energy X-ray Absorptiometry).
|
||||||
|
|
||||||
|
### Ключевые возможности
|
||||||
|
|
||||||
|
| Возможность | Описание |
|
||||||
|
|-------------|----------|
|
||||||
|
| 📥 Загрузка DICOM | Работа с медицинскими изображениями в стандартном формате |
|
||||||
|
| 🔍 Анализ качества | Оценка по критериям: позиционирование, артефакты, контраст |
|
||||||
|
| 🏷️ Классификация | Бинарная классификация: OK / Нарушение |
|
||||||
|
| 📊 Экспорт | Вывод результатов в XLSX/CSV формате |
|
||||||
|
| 🌐 REST API | Интеграция с внешними системами |
|
||||||
|
| 🖥️ Web-интерфейс | Удобная визуализация и загрузка |
|
||||||
|
|
||||||
|
### Типы анализируемых исследований
|
||||||
|
|
||||||
|
- **Позвоночник** (spine)
|
||||||
|
- **Бедро правое** (hip_right)
|
||||||
|
- **Бедро левое** (hip_left)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 2. Проблема и актуальность
|
||||||
|
|
||||||
|
### Проблема
|
||||||
|
|
||||||
|
При проведении денситометрии (DXA) критически важно соблюдать стандарты качества исследования:
|
||||||
|
|
||||||
|
- ❌ Неправильное позиционирование пациента
|
||||||
|
- ❌ Артефакты движения
|
||||||
|
- ❌ Недостаточный контраст
|
||||||
|
- ❌ Неправильная укладка
|
||||||
|
|
||||||
|
Нарушение этих критериев приводит к:
|
||||||
|
- ❌ Неправильной диагностике остеопороза
|
||||||
|
- ❌ Повторным исследованиям
|
||||||
|
- ❌ Дополнительной лучевой нагрузке
|
||||||
|
- ❌ Потере времени и ресурсов
|
||||||
|
|
||||||
|
### Решение
|
||||||
|
|
||||||
|
Автоматизированная система оценки качества, которая:
|
||||||
|
- ✅ Мгновенно анализирует каждое изображение
|
||||||
|
- ✅ Выявляет конкретные нарушения
|
||||||
|
- ✅ Обеспечивает стандартизацию контроля качества
|
||||||
|
- ✅ Интегрируется в рабочий процесс рентгенолога
|
||||||
|
|
||||||
|
### Бизнес-ценность
|
||||||
|
|
||||||
|
```
|
||||||
|
Экономия времени рентгенолога: ~30-60 секунд на исследование
|
||||||
|
Снижение количества повторных исследований: до 15-20%
|
||||||
|
Повышение качества диагностики: стандартизированный контроль
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 3. Техническая архитектура
|
||||||
|
|
||||||
|
### Общая схема
|
||||||
|
|
||||||
|
```
|
||||||
|
┌─────────────┐ ┌──────────────┐ ┌─────────────┐
|
||||||
|
│ DICOM │────▶│ FastAPI │────▶│ PyTorch │
|
||||||
|
│ Файлы │ │ Server │ │ Model │
|
||||||
|
└─────────────┘ └──────────────┘ └─────────────┘
|
||||||
|
│
|
||||||
|
▼
|
||||||
|
┌──────────────┐
|
||||||
|
│ XLSX/CSV │
|
||||||
|
│ Output │
|
||||||
|
└──────────────┘
|
||||||
|
```
|
||||||
|
|
||||||
|
### Компоненты системы
|
||||||
|
|
||||||
|
| Компонент | Технология | Назначение |
|
||||||
|
|-----------|------------|------------|
|
||||||
|
| Backend | Python 3.10, FastAPI | REST API, веб-сервер |
|
||||||
|
| ML Framework | PyTorch 2.1 | Глубокое обучение |
|
||||||
|
| Модель | ResNet18 (ImageNet) | Классификация изображений |
|
||||||
|
| Обработка изображений | PIL, OpenCV | Предобработка DICOM |
|
||||||
|
| Данные | pandas, openpyxl | Работа с Excel |
|
||||||
|
| Контейнеризация | Docker | Деплой |
|
||||||
|
|
||||||
|
### Структура проекта
|
||||||
|
|
||||||
|
```
|
||||||
|
bone_2026/
|
||||||
|
├── src/
|
||||||
|
│ ├── main.py # FastAPI приложение
|
||||||
|
│ ├── dxa/
|
||||||
|
│ │ ├── dataset.py # Загрузка данных
|
||||||
|
│ │ ├── model.py # Архитектура модели
|
||||||
|
│ │ ├── train.py # Обучение
|
||||||
|
│ │ └── inference.py # Инференс
|
||||||
|
│ └── api/ # Эндпоинты
|
||||||
|
├── models/
|
||||||
|
│ └── dxa_model.pth # Обученная модель
|
||||||
|
├── dataset_hack/ # Датасет
|
||||||
|
└── requirements.txt # Зависимости
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 4. Датасет и разметка
|
||||||
|
|
||||||
|
### Источник данных
|
||||||
|
|
||||||
|
- **100 исследований** из реальной клинической практики
|
||||||
|
- **499 DICOM файлов**
|
||||||
|
- **Аннотации** в формате Excel
|
||||||
|
|
||||||
|
### Типы анатомических регионов
|
||||||
|
|
||||||
|
| Регион | Критерии оценки |
|
||||||
|
|--------|-----------------|
|
||||||
|
| Позвоночник | Укладка, Ось, Артефакты |
|
||||||
|
| Бедро (правое/левое) | Позиция, ROI |
|
||||||
|
|
||||||
|
### Классы качества
|
||||||
|
|
||||||
|
| Класс | Описание | Примерное соотношение |
|
||||||
|
|-------|----------|----------------------|
|
||||||
|
| 0 | OK (качество соответствует стандартам) | ~65% |
|
||||||
|
| 1 | Нарушение (выявлены проблемы) | ~35% |
|
||||||
|
|
||||||
|
### Пример разметки
|
||||||
|
|
||||||
|
```
|
||||||
|
Study UID | Позвоночник_итого | Бедро_прав_итого | Бедро_лев_итого | Класс
|
||||||
|
----------|-------------------|------------------|-----------------|------
|
||||||
|
1.2.840...| 0 | 1 | 0 | 1
|
||||||
|
1.2.840...| 0 | 0 | 0 | 0
|
||||||
|
```
|
||||||
|
|
||||||
|
### Разделение данных
|
||||||
|
|
||||||
|
- **Обучение**: ~80% (1146 samples)
|
||||||
|
- **Валидация**: ~20% (287 samples)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 5. ML-модель
|
||||||
|
|
||||||
|
### Архитектура
|
||||||
|
|
||||||
|
```
|
||||||
|
DXAQualityClassifier (ResNet18)
|
||||||
|
├── Backbone: ResNet18 (pretrained on ImageNet)
|
||||||
|
│ └── Feature extraction: 512 dimensions
|
||||||
|
└── Classifier Head:
|
||||||
|
├── Dropout (0.3)
|
||||||
|
├── Linear (512 → 256)
|
||||||
|
├── ReLU
|
||||||
|
├── Dropout (0.3)
|
||||||
|
└── Linear (256 → 2)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Входные данные
|
||||||
|
|
||||||
|
- **Размер изображения**: 224 × 224 пикселя
|
||||||
|
- **Формат**: RGB (3 канала)
|
||||||
|
- **Нормализация**: ImageNet stats
|
||||||
|
|
||||||
|
### Выходные данные
|
||||||
|
|
||||||
|
- **Класс 0**: Качество OK
|
||||||
|
- **Класс 1**: Нарушение качества
|
||||||
|
- **Confidence**: Вероятность предсказания
|
||||||
|
|
||||||
|
### Гиперпараметры обучения
|
||||||
|
|
||||||
|
| Параметр | Значение |
|
||||||
|
|----------|----------|
|
||||||
|
| Backbone | ResNet18 |
|
||||||
|
| Input size | 224 × 224 |
|
||||||
|
| Batch size | 8-16 |
|
||||||
|
| Learning rate | 1e-4 |
|
||||||
|
| Optimizer | AdamW |
|
||||||
|
| Scheduler | ReduceLROnPlateau |
|
||||||
|
| Dropout | 0.3 |
|
||||||
|
|
||||||
|
### Пример команды обучения
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python src/dxa/train.py \
|
||||||
|
--epochs 10 \
|
||||||
|
--batch-size 16 \
|
||||||
|
--backbone resnet18 \
|
||||||
|
--data-root dataset_hack
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 6. API и интеграция
|
||||||
|
|
||||||
|
### REST Endpoints
|
||||||
|
|
||||||
|
| Метод | Эндпоинт | Описание |
|
||||||
|
|-------|----------|----------|
|
||||||
|
| GET | `/` | Главная страница (web UI) |
|
||||||
|
| GET | `/api/v1/health` | Проверка статуса |
|
||||||
|
| POST | `/api/v1/analyze` | Анализ одного DICOM |
|
||||||
|
| POST | `/api/v1/batch` | Пакетный анализ |
|
||||||
|
| POST | `/api/v1/export` | Анализ + экспорт XLSX |
|
||||||
|
|
||||||
|
### Пример запроса (cURL)
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl -X POST "http://localhost:8000/api/v1/analyze" \
|
||||||
|
-H "accept: application/json" \
|
||||||
|
-H "Content-Type: multipart/form-data" \
|
||||||
|
-F "file=@/path/to/image.dcm"
|
||||||
|
```
|
||||||
|
|
||||||
|
### Пример ответа
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"study_uid": "1.2.840.113619.2.110.512719.20250403090444",
|
||||||
|
"image_uid": "1.2.3.4.5.6.7.8.9",
|
||||||
|
"anatomical_region": "spine",
|
||||||
|
"quality_class": 0,
|
||||||
|
"quality_label": "OK",
|
||||||
|
"confidence": 0.9234,
|
||||||
|
"processing_status": "Success"
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Формат вывода (XLSX/CSV)
|
||||||
|
|
||||||
|
| Колонка | Описание |
|
||||||
|
|---------|----------|
|
||||||
|
| path_to_study | Путь к исследованию |
|
||||||
|
| study_uid | StudyInstanceUID |
|
||||||
|
| image_uid | SOPInstanceUID |
|
||||||
|
| anatomical_region | Регион (spine/hip) |
|
||||||
|
| quality_class | Класс (0/1) |
|
||||||
|
| violation_type | Тип нарушения |
|
||||||
|
| processing_status | Статус обработки |
|
||||||
|
| time_of_processing | Время обработки (сек) |
|
||||||
|
|
||||||
|
### Web-интерфейс
|
||||||
|
|
||||||
|
- 📤 Drag-and-drop загрузка
|
||||||
|
- 🔍 Автоматический анализ
|
||||||
|
- 📊 Визуализация результатов
|
||||||
|
- 📥 Экспорт в Excel
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 7. Результаты
|
||||||
|
|
||||||
|
### Метрики на валидации
|
||||||
|
|
||||||
|
| Метрика | Значение |
|
||||||
|
|---------|----------|
|
||||||
|
| Accuracy | ~84% |
|
||||||
|
| Precision | ~0.35 |
|
||||||
|
| Recall | ~0.22 |
|
||||||
|
| **F1 Score** | **~0.27** |
|
||||||
|
|
||||||
|
### Анализ результатов
|
||||||
|
|
||||||
|
```
|
||||||
|
Распределение классов в валидации:
|
||||||
|
├── OK (класс 0): ████████████████████ 65%
|
||||||
|
└── Нарушение (класс 1): ██████████ 35%
|
||||||
|
|
||||||
|
Примечание: F1 скорее всего ниже из-за
|
||||||
|
несбалансированности классов и ограниченного
|
||||||
|
количества данных (100 исследований)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Время обработки
|
||||||
|
|
||||||
|
- **Одно изображение**: ~50-100 мс
|
||||||
|
- **Пакетная обработка**: зависит от размера батча
|
||||||
|
|
||||||
|
### Ограничения текущей версии
|
||||||
|
|
||||||
|
- ⚠️ Ограниченный объём данных (100 исследований)
|
||||||
|
- ⚠️ Дисбаланс классов
|
||||||
|
- ⚠️ Только бинарная классификация
|
||||||
|
- ⚠️ Требуется больше эпох обучения
|
||||||
|
|
||||||
|
### Рекомендации для улучшения
|
||||||
|
|
||||||
|
1. **Увеличение датасета**: 500+ исследований
|
||||||
|
2. **Балансировка классов**: weighted loss, oversampling
|
||||||
|
3. **Аугментация**: rotation, flip, brightness
|
||||||
|
4. **Архитектура**: EfficientNet, ViT
|
||||||
|
5. **Многоклассовая классификация**: типы нарушений
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 8. Демонстрация
|
||||||
|
|
||||||
|
### Сценарий использования
|
||||||
|
|
||||||
|
1. **Врач загружает DICOM файл** через web-интерфейс или API
|
||||||
|
2. **Система автоматически:**
|
||||||
|
- Определяет анатомический регион
|
||||||
|
- Анализирует качество изображения
|
||||||
|
- Классифицирует: OK / Нарушение
|
||||||
|
3. **Результат:**
|
||||||
|
- Визуализация в интерфейсе
|
||||||
|
- Excel-отчёт для дальнейшего анализа
|
||||||
|
|
||||||
|
### Пример работы
|
||||||
|
|
||||||
|
```
|
||||||
|
Вход: DICOM файл исследования позвоночника
|
||||||
|
↓
|
||||||
|
Обработка: ResNet18 классификатор
|
||||||
|
↓
|
||||||
|
Выход: {
|
||||||
|
"anatomical_region": "spine",
|
||||||
|
"quality_class": 0,
|
||||||
|
"confidence": 0.89
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### Screenshots (доступны в интерфейсе)
|
||||||
|
|
||||||
|
- Загрузка файла
|
||||||
|
- Результат анализа
|
||||||
|
- Экспорт отчёта
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 9. Перспективы развития
|
||||||
|
|
||||||
|
### Краткосрочные улучшения
|
||||||
|
|
||||||
|
| Направление | Описание |
|
||||||
|
|-------------|----------|
|
||||||
|
| Увеличение данных | Добавить 400+ исследований |
|
||||||
|
| Классовый баланс | Weighted loss, SMOTE |
|
||||||
|
| Аугментация | Rotation, flip, brightness |
|
||||||
|
| Типы нарушений | Multiclass: позиция, артефакты, контраст |
|
||||||
|
|
||||||
|
### Среднесрочные цели
|
||||||
|
|
||||||
|
| Направление | Описание |
|
||||||
|
|-------------|----------|
|
||||||
|
| Сегментация | Выделение позвонков, бедра |
|
||||||
|
| Многозадачность | Классификация + сегментация |
|
||||||
|
| AutoML | Оптимальная архитектура |
|
||||||
|
| Мониторинг | Логирование, метрики |
|
||||||
|
|
||||||
|
### Долгосрочное видение
|
||||||
|
|
||||||
|
- 🏥 **Интеграция с PACS**
|
||||||
|
- 📱 **Мобильное приложение**
|
||||||
|
- 🤖 **Real-time анализ**
|
||||||
|
- 🌐 **Облачное решение (SaaS)**
|
||||||
|
- 📊 **Дашборд для радиологов**
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Контакты
|
||||||
|
|
||||||
|
| | |
|
||||||
|
|---|---|
|
||||||
|
| **Автор** | Грачев Денис |
|
||||||
|
| **Email** | denis@example.com |
|
||||||
|
| **Telegram** | @oxydencher |
|
||||||
|
| **GitHub** | github.com/yourusername |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
*Проект разработан в рамках хакатона по медицинскому ИИ*
|
||||||
|
|
@ -0,0 +1,111 @@
|
||||||
|
# Задачи для следующей сессии
|
||||||
|
|
||||||
|
## Сравнение текущей реализации с требованиями
|
||||||
|
|
||||||
|
### ✅ РЕАЛИЗОВАНО
|
||||||
|
|
||||||
|
| Требование | Статус | Файл |
|
||||||
|
|------------|--------|------|
|
||||||
|
| Бинарная классификация OK/Нарушение | ✅ Готово | src/dxa/model.py |
|
||||||
|
| Определение анатомической области | ✅ Готово | src/dxa/inference.py |
|
||||||
|
| Экспорт в XLSX/CSV | ✅ Готово | src/main.py /api/v1/export |
|
||||||
|
| API для пакетной обработки | ✅ Готово | src/main.py /api/v1/batch |
|
||||||
|
| Docker контейнеризация | ✅ Готово | Dockerfile |
|
||||||
|
| Web-интерфейс | ✅ Готово | src/api/static/ |
|
||||||
|
| Определение типа нарушения | ✅ Частично | src/quality/detailed_assessment.py |
|
||||||
|
| Детальные метрики качества | ✅ Готово | src/quality/detailed_assessment.py |
|
||||||
|
| Проверка движения (motion) | ✅ Готово | detect_motion_blur() |
|
||||||
|
| Проверка артефактов | ✅ Готово | detect_artifacts() |
|
||||||
|
| DICOM SR формат | ✅ Готово | /api/v1/analyze/sr |
|
||||||
|
| Визуализация маски | ✅ Готово | include_visualization=true |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Мультимодельная архитектура
|
||||||
|
|
||||||
|
См. документацию: `docs/multi_model_architecture.md`
|
||||||
|
|
||||||
|
```
|
||||||
|
Pipeline:
|
||||||
|
1. Region Detector → 2. Segmentator → 3. Quality Classifier → 4. Violation Type → 5. Aggregator
|
||||||
|
```
|
||||||
|
|
||||||
|
### Модели:
|
||||||
|
|
||||||
|
| Модель | Назначение | Файл |
|
||||||
|
|--------|------------|------|
|
||||||
|
| Region Detector | Spine/Hip определение | `src/models/region_detector.py` |
|
||||||
|
| Segmentator | Сегментация костей | `src/models/segmentation/` |
|
||||||
|
| Quality Classifier | OK/Violation бинарный | `src/models/classification/quality.py` |
|
||||||
|
| Violation Classifier | Типы нарушений (7+) | `src/models/classification/violation.py` |
|
||||||
|
| Artifact Detector | Детекция артефактов | `src/models/artifacts/detector.py` |
|
||||||
|
| Grad-CAM | Heatmap внимания | `src/models/visualization/gradcam.py` |
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## ПЛАН РЕАЛИЗАЦИИ
|
||||||
|
|
||||||
|
### Фаза 1: Архитектура и модели (высокий приоритет)
|
||||||
|
|
||||||
|
1. **Создать структуру моделей**
|
||||||
|
- `src/models/` - директория для всех моделей
|
||||||
|
- `src/pipeline/` - основной пайплайн
|
||||||
|
- `src/training/` - скрипты обучения
|
||||||
|
|
||||||
|
2. **Region Detector**
|
||||||
|
- Переиспользовать текущую логику из inference.py
|
||||||
|
- Обучить классификатор на 4 класса
|
||||||
|
|
||||||
|
3. **Segmentation Models**
|
||||||
|
- Spine Segmentator (U-Net)
|
||||||
|
- Hip Segmentator (U-Net)
|
||||||
|
|
||||||
|
4. **Quality Classifier (бинарный)**
|
||||||
|
- ResNet18 с attention
|
||||||
|
- Weighted loss для балансировки
|
||||||
|
|
||||||
|
5. **Violation Type Classifier (многоклассовый)**
|
||||||
|
- 7 классов для spine
|
||||||
|
- 7 классов для hip
|
||||||
|
- Multi-task learning
|
||||||
|
|
||||||
|
### Фаза 2: Пайплайн и интеграция
|
||||||
|
|
||||||
|
6. **Pipeline класс**
|
||||||
|
- Объединение всех моделей
|
||||||
|
- Queue для асинхронной обработки
|
||||||
|
|
||||||
|
7. **API интеграция**
|
||||||
|
- Обновить endpoints для использования пайплайна
|
||||||
|
- Batch processing
|
||||||
|
|
||||||
|
### Фаза 3: Визуализация и метрики
|
||||||
|
|
||||||
|
8. **Grad-CAM Heatmap**
|
||||||
|
- Извлечение attention карт
|
||||||
|
- Наложение на изображение
|
||||||
|
|
||||||
|
9. **Метрики**
|
||||||
|
- ROC-AUC, PR-AUC
|
||||||
|
- Macro-F1 с доверительными интервалами
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## Текущие ограничения
|
||||||
|
|
||||||
|
1. Модель обучена на бинарную классификацию - нужно переобучение
|
||||||
|
2. F1 ~0.27 - низкий из-за дисбаланса классов
|
||||||
|
3. Датасет ~100 исследований - нужно 500+
|
||||||
|
4. Сегментация использует простой threshold - нужно дообучить модель
|
||||||
|
5. Heatmap не реализовано - требует дообучения
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## ЗАПУСК ДЛЯ ТЕСТИРОВАНИЯ
|
||||||
|
|
||||||
|
```bash
|
||||||
|
cd /Users/denis/workspace/bone_2026
|
||||||
|
python -m uvicorn src.main:app --host 0.0.0.0 --port 8000
|
||||||
|
```
|
||||||
|
|
||||||
|
Открыть http://localhost:8000
|
||||||
|
|
@ -0,0 +1,32 @@
|
||||||
|
"""
|
||||||
|
DXA Quality Assessment Models
|
||||||
|
"""
|
||||||
|
from .base import BaseDXAModel, BaseClassifier, BaseSegmentator
|
||||||
|
from .region_detector import RegionDetector, create_region_detector, determine_region_from_image
|
||||||
|
from .classification import QualityClassifier, ViolationTypeClassifier
|
||||||
|
from .segmentation import DXASegmenter, create_segmenter
|
||||||
|
from .visualization import GradCAMExtractor, visualize_attention
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
# Base
|
||||||
|
'BaseDXAModel',
|
||||||
|
'BaseClassifier',
|
||||||
|
'BaseSegmentator',
|
||||||
|
|
||||||
|
# Region
|
||||||
|
'RegionDetector',
|
||||||
|
'create_region_detector',
|
||||||
|
'determine_region_from_image',
|
||||||
|
|
||||||
|
# Classification
|
||||||
|
'QualityClassifier',
|
||||||
|
'ViolationTypeClassifier',
|
||||||
|
|
||||||
|
# Segmentation
|
||||||
|
'DXASegmenter',
|
||||||
|
'create_segmenter',
|
||||||
|
|
||||||
|
# Visualization
|
||||||
|
'GradCAMExtractor',
|
||||||
|
'visualize_attention'
|
||||||
|
]
|
||||||
|
|
@ -0,0 +1,92 @@
|
||||||
|
"""
|
||||||
|
Base model class for all DXA models
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from typing import Dict, Any, Optional
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
|
||||||
|
|
||||||
|
class BaseDXAModel(nn.Module, ABC):
|
||||||
|
"""Base class for all DXA models"""
|
||||||
|
|
||||||
|
def __init__(self, config: Optional[Dict] = None):
|
||||||
|
super().__init__()
|
||||||
|
self.config = config or {}
|
||||||
|
self.device = torch.device('cpu')
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Forward pass"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
def extract_features(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Extract features without classification"""
|
||||||
|
return self.forward(x)
|
||||||
|
|
||||||
|
def to_device(self, device: str = 'cpu'):
|
||||||
|
"""Move model to device"""
|
||||||
|
self.device = torch.device(device)
|
||||||
|
self.to(self.device)
|
||||||
|
return self
|
||||||
|
|
||||||
|
def get_info(self) -> Dict[str, Any]:
|
||||||
|
"""Get model info"""
|
||||||
|
return {
|
||||||
|
'name': self.__class__.__name__,
|
||||||
|
'device': str(self.device),
|
||||||
|
'config': self.config
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class BaseClassifier(BaseDXAModel):
|
||||||
|
"""Base classifier"""
|
||||||
|
|
||||||
|
def __init__(self, num_classes: int = 2, **kwargs):
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
self.num_classes = num_classes
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def predict(self, x: torch.Tensor) -> Dict[str, Any]:
|
||||||
|
"""Predict on input"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
def predict_proba(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Get probabilities"""
|
||||||
|
self.eval()
|
||||||
|
with torch.no_grad():
|
||||||
|
logits = self.forward(x)
|
||||||
|
return torch.softmax(logits, dim=1)
|
||||||
|
|
||||||
|
def predict_class(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Get class predictions"""
|
||||||
|
probs = self.predict_proba(x)
|
||||||
|
return probs.argmax(dim=1)
|
||||||
|
|
||||||
|
|
||||||
|
class BaseSegmentator(BaseDXAModel):
|
||||||
|
"""Base segmentator"""
|
||||||
|
|
||||||
|
def __init__(self, in_channels: int = 3, out_classes: int = 2, **kwargs):
|
||||||
|
super().__init__(**kwargs)
|
||||||
|
self.in_channels = in_channels
|
||||||
|
self.out_classes = out_classes
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def predict(self, x: torch.Tensor) -> Dict[str, Any]:
|
||||||
|
"""Predict on input"""
|
||||||
|
pass
|
||||||
|
|
||||||
|
def predict_mask(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Get segmentation mask"""
|
||||||
|
self.eval()
|
||||||
|
with torch.no_grad():
|
||||||
|
logits = self.forward(x)
|
||||||
|
return logits.argmax(dim=1)
|
||||||
|
|
||||||
|
def predict_proba(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Get class probabilities"""
|
||||||
|
self.eval()
|
||||||
|
with torch.no_grad():
|
||||||
|
logits = self.forward(x)
|
||||||
|
return torch.softmax(logits, dim=1)
|
||||||
|
|
@ -0,0 +1,21 @@
|
||||||
|
"""
|
||||||
|
Classification models
|
||||||
|
"""
|
||||||
|
from .quality_classifier import QualityClassifier, create_quality_classifier
|
||||||
|
from .violation_classifier import (
|
||||||
|
ViolationTypeClassifier,
|
||||||
|
create_violation_classifier,
|
||||||
|
SPINE_VIOLATIONS,
|
||||||
|
HIP_VIOLATIONS,
|
||||||
|
VIOLATION_DESCRIPTIONS
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
'QualityClassifier',
|
||||||
|
'create_quality_classifier',
|
||||||
|
'ViolationTypeClassifier',
|
||||||
|
'create_violation_classifier',
|
||||||
|
'SPINE_VIOLATIONS',
|
||||||
|
'HIP_VIOLATIONS',
|
||||||
|
'VIOLATION_DESCRIPTIONS'
|
||||||
|
]
|
||||||
|
|
@ -0,0 +1,224 @@
|
||||||
|
"""
|
||||||
|
Quality Classifier - бинарный классификатор OK/Violation
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torchvision.models as models
|
||||||
|
from typing import Dict, Tuple, Optional
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
|
||||||
|
class QualityClassifier(nn.Module):
|
||||||
|
"""
|
||||||
|
Бинарный классификатор качества DXA исследования.
|
||||||
|
Выход: 0 (OK) или 1 (Violation)
|
||||||
|
"""
|
||||||
|
|
||||||
|
QUALITY_LABELS = {0: 'OK', 1: 'Violation'}
|
||||||
|
|
||||||
|
def __init__(self,
|
||||||
|
backbone: str = 'resnet18',
|
||||||
|
pretrained: bool = True,
|
||||||
|
dropout: float = 0.4,
|
||||||
|
use_attention: bool = False):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.backbone_name = backbone
|
||||||
|
self.use_attention = use_attention
|
||||||
|
|
||||||
|
# Encoder для изображения
|
||||||
|
if backbone == 'resnet18':
|
||||||
|
self.image_encoder = models.resnet18(
|
||||||
|
weights='IMAGENET1K_V1' if pretrained else None
|
||||||
|
)
|
||||||
|
self.image_encoder.fc = nn.Identity()
|
||||||
|
image_features = 512
|
||||||
|
elif backbone == 'resnet34':
|
||||||
|
self.image_encoder = models.resnet34(
|
||||||
|
weights='IMAGENET1K_V1' if pretrained else None
|
||||||
|
)
|
||||||
|
self.image_encoder.fc = nn.Identity()
|
||||||
|
image_features = 512
|
||||||
|
elif backbone == 'efficientnet_b0':
|
||||||
|
self.image_encoder = models.efficientnet_b0(
|
||||||
|
weights='IMAGENET1K_V1' if pretrained else None
|
||||||
|
)
|
||||||
|
self.image_encoder.classifier = nn.Identity()
|
||||||
|
image_features = 1280
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown backbone: {backbone}")
|
||||||
|
|
||||||
|
# Encoder для маски (опционально)
|
||||||
|
if use_attention:
|
||||||
|
self.mask_encoder = nn.Sequential(
|
||||||
|
nn.Conv2d(1, 32, 3, padding=1),
|
||||||
|
nn.BatchNorm2d(32),
|
||||||
|
nn.ReLU(inplace=True),
|
||||||
|
nn.Conv2d(32, 64, 3, stride=2, padding=1),
|
||||||
|
nn.BatchNorm2d(64),
|
||||||
|
nn.ReLU(inplace=True),
|
||||||
|
nn.AdaptiveAvgPool2d(1)
|
||||||
|
)
|
||||||
|
combined_features = image_features + 64
|
||||||
|
else:
|
||||||
|
self.mask_encoder = None
|
||||||
|
combined_features = image_features
|
||||||
|
|
||||||
|
# Классификатор
|
||||||
|
self.classifier = nn.Sequential(
|
||||||
|
nn.Dropout(dropout),
|
||||||
|
nn.Linear(combined_features, 256),
|
||||||
|
nn.ReLU(inplace=True),
|
||||||
|
nn.Dropout(dropout),
|
||||||
|
nn.Linear(256, 2) # OK / Violation
|
||||||
|
)
|
||||||
|
|
||||||
|
self.image_features = image_features
|
||||||
|
|
||||||
|
def forward(self, image: torch.Tensor, mask: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
Forward pass
|
||||||
|
|
||||||
|
Args:
|
||||||
|
image: [B, 3, H, W] RGB изображение
|
||||||
|
mask: [B, 1, H, W] опциональная маска сегментации
|
||||||
|
"""
|
||||||
|
# Извлечение признаков из изображения
|
||||||
|
img_features = self.image_encoder(image)
|
||||||
|
|
||||||
|
# Если есть маска и используется attention
|
||||||
|
if mask is not None and self.use_attention and self.mask_encoder is not None:
|
||||||
|
mask_features = self.mask_encoder(mask).squeeze(-1).squeeze(-1)
|
||||||
|
combined = torch.cat([img_features, mask_features], dim=1)
|
||||||
|
else:
|
||||||
|
combined = img_features
|
||||||
|
|
||||||
|
return self.classifier(combined)
|
||||||
|
|
||||||
|
def predict(self,
|
||||||
|
image: torch.Tensor,
|
||||||
|
mask: Optional[torch.Tensor] = None) -> Dict:
|
||||||
|
"""
|
||||||
|
Предсказание качества.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict: predicted_class, label, confidence, probabilities
|
||||||
|
"""
|
||||||
|
self.eval()
|
||||||
|
with torch.no_grad():
|
||||||
|
logits = self.forward(image, mask)
|
||||||
|
probs = torch.softmax(logits, dim=1)
|
||||||
|
conf, pred = probs.max(dim=1)
|
||||||
|
|
||||||
|
results = []
|
||||||
|
for i in range(image.shape[0]):
|
||||||
|
pred_class = pred[i].item()
|
||||||
|
results.append({
|
||||||
|
'predicted_class': pred_class,
|
||||||
|
'label': self.QUALITY_LABELS[pred_class],
|
||||||
|
'confidence': conf[i].item(),
|
||||||
|
'probabilities': {
|
||||||
|
'OK': probs[i, 0].item(),
|
||||||
|
'Violation': probs[i, 1].item()
|
||||||
|
},
|
||||||
|
'logits': logits[i].cpu().numpy().tolist()
|
||||||
|
})
|
||||||
|
|
||||||
|
return results[0] if len(results) == 1 else results
|
||||||
|
|
||||||
|
def extract_features(self,
|
||||||
|
image: torch.Tensor,
|
||||||
|
mask: Optional[torch.Tensor] = None) -> torch.Tensor:
|
||||||
|
"""Извлечение признаков"""
|
||||||
|
img_features = self.image_encoder(image)
|
||||||
|
|
||||||
|
if mask is not None and self.use_attention and self.mask_encoder is not None:
|
||||||
|
mask_features = self.mask_encoder(mask).squeeze(-1).squeeze(-1)
|
||||||
|
return torch.cat([img_features, mask_features], dim=1)
|
||||||
|
|
||||||
|
return img_features
|
||||||
|
|
||||||
|
def get_info(self) -> Dict:
|
||||||
|
return {
|
||||||
|
'name': 'QualityClassifier',
|
||||||
|
'backbone': self.backbone_name,
|
||||||
|
'num_classes': 2,
|
||||||
|
'use_attention': self.use_attention,
|
||||||
|
'image_features': self.image_features
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class QualityClassifierMultiRegion(nn.Module):
|
||||||
|
"""
|
||||||
|
Мультирегиональный классификатор качества.
|
||||||
|
Отдельные головы для spine и hip.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self,
|
||||||
|
backbone: str = 'resnet18',
|
||||||
|
pretrained: bool = True,
|
||||||
|
dropout: float = 0.4):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
# Общий encoder
|
||||||
|
if backbone == 'resnet18':
|
||||||
|
self.encoder = models.resnet18(
|
||||||
|
weights='IMAGENET1K_V1' if pretrained else None
|
||||||
|
)
|
||||||
|
self.encoder.fc = nn.Identity()
|
||||||
|
feature_dim = 512
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown backbone: {backbone}")
|
||||||
|
|
||||||
|
# Отдельные головы для каждого региона
|
||||||
|
self.spine_head = nn.Sequential(
|
||||||
|
nn.Dropout(dropout),
|
||||||
|
nn.Linear(feature_dim, 128),
|
||||||
|
nn.ReLU(),
|
||||||
|
nn.Linear(128, 2)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.hip_head = nn.Sequential(
|
||||||
|
nn.Dropout(dropout),
|
||||||
|
nn.Linear(feature_dim, 128),
|
||||||
|
nn.ReLU(),
|
||||||
|
nn.Linear(128, 2)
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor, region: str = 'spine') -> torch.Tensor:
|
||||||
|
features = self.encoder(x)
|
||||||
|
|
||||||
|
if region == 'spine':
|
||||||
|
return self.spine_head(features)
|
||||||
|
else:
|
||||||
|
return self.hip_head(features)
|
||||||
|
|
||||||
|
def predict(self, image: torch.Tensor, region: str = 'spine') -> Dict:
|
||||||
|
self.eval()
|
||||||
|
with torch.no_grad():
|
||||||
|
logits = self.forward(image, region)
|
||||||
|
probs = torch.softmax(logits, dim=1)
|
||||||
|
conf, pred = probs.max(dim=1)
|
||||||
|
|
||||||
|
return {
|
||||||
|
'predicted_class': pred.item(),
|
||||||
|
'label': QualityClassifier.QUALITY_LABELS[pred.item()],
|
||||||
|
'confidence': conf.item(),
|
||||||
|
'region': region
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def create_quality_classifier(
|
||||||
|
backbone: str = 'resnet18',
|
||||||
|
pretrained: bool = True,
|
||||||
|
use_attention: bool = False,
|
||||||
|
device: str = 'cpu'
|
||||||
|
) -> QualityClassifier:
|
||||||
|
"""Создание классификатора качества"""
|
||||||
|
model = QualityClassifier(
|
||||||
|
backbone=backbone,
|
||||||
|
pretrained=pretrained,
|
||||||
|
use_attention=use_attention
|
||||||
|
)
|
||||||
|
model.to(device)
|
||||||
|
return model
|
||||||
|
|
@ -0,0 +1,237 @@
|
||||||
|
"""
|
||||||
|
Violation Type Classifier - многоклассовый классификатор типа нарушения
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torchvision.models as models
|
||||||
|
from typing import Dict, List, Optional
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
|
||||||
|
# Типы нарушений для позвоночника
|
||||||
|
SPINE_VIOLATIONS = [
|
||||||
|
'correct',
|
||||||
|
'position_error', # Неправильная укладка
|
||||||
|
'artifact_motion', # Артефакт движения
|
||||||
|
'artifact_other', # Другие артефакты
|
||||||
|
'labeling_error', # Ошибка разметки
|
||||||
|
'incomplete_view', # Неполный вид
|
||||||
|
'roi_error' # Ошибка ROI
|
||||||
|
]
|
||||||
|
|
||||||
|
# Типы нарушений для бедра
|
||||||
|
HIP_VIOLATIONS = [
|
||||||
|
'correct',
|
||||||
|
'position_error',
|
||||||
|
'rotation', # Нарушение ротации
|
||||||
|
'artifact_motion',
|
||||||
|
'artifact_other',
|
||||||
|
'roi_error',
|
||||||
|
'incomplete_view'
|
||||||
|
]
|
||||||
|
|
||||||
|
# Русские описания
|
||||||
|
VIOLATION_DESCRIPTIONS = {
|
||||||
|
'correct': 'Качество соответствует норме',
|
||||||
|
'position_error': 'Неправильное положение пациента',
|
||||||
|
'artifact_motion': 'Артефакт движения (размытие)',
|
||||||
|
'artifact_other': 'Другие артефакты',
|
||||||
|
'labeling_error': 'Ошибка разметки',
|
||||||
|
'incomplete_view': 'Неполный вид анатомической области',
|
||||||
|
'roi_error': 'Неправильное положение ROI',
|
||||||
|
'rotation': 'Нарушение ротации'
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class ViolationTypeClassifier(nn.Module):
|
||||||
|
"""
|
||||||
|
Многоклассовый классификатор типа нарушения.
|
||||||
|
Поддерживает разные классификаторы для spine и hip.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self,
|
||||||
|
backbone: str = 'resnet18',
|
||||||
|
num_classes_spine: int = 7,
|
||||||
|
num_classes_hip: int = 7,
|
||||||
|
pretrained: bool = True,
|
||||||
|
dropout: float = 0.4,
|
||||||
|
use_multi_task: bool = True):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.backbone_name = backbone
|
||||||
|
self.num_classes_spine = num_classes_spine
|
||||||
|
self.num_classes_hip = num_classes_hip
|
||||||
|
self.use_multi_task = use_multi_task
|
||||||
|
|
||||||
|
# Encoder
|
||||||
|
if backbone == 'resnet18':
|
||||||
|
self.encoder = models.resnet18(
|
||||||
|
weights='IMAGENET1K_V1' if pretrained else None
|
||||||
|
)
|
||||||
|
self.encoder.fc = nn.Identity()
|
||||||
|
feature_dim = 512
|
||||||
|
elif backbone == 'resnet34':
|
||||||
|
self.encoder = models.resnet34(
|
||||||
|
weights='IMAGENET1K_V1' if pretrained else None
|
||||||
|
)
|
||||||
|
self.encoder.fc = nn.Identity()
|
||||||
|
feature_dim = 512
|
||||||
|
elif backbone == 'efficientnet_b0':
|
||||||
|
self.encoder = models.efficientnet_b0(
|
||||||
|
weights='IMAGENET1K_V1' if pretrained else None
|
||||||
|
)
|
||||||
|
self.encoder.classifier = nn.Identity()
|
||||||
|
feature_dim = 1280
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown backbone: {backbone}")
|
||||||
|
|
||||||
|
if use_multi_task:
|
||||||
|
# Общая голова с регионом
|
||||||
|
self.head = nn.Sequential(
|
||||||
|
nn.Dropout(dropout),
|
||||||
|
nn.Linear(feature_dim, 256),
|
||||||
|
nn.ReLU(),
|
||||||
|
nn.Dropout(dropout),
|
||||||
|
nn.Linear(256, 14) # 7 + 7 = 14 классов
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Отдельные головы
|
||||||
|
self.spine_head = nn.Sequential(
|
||||||
|
nn.Dropout(dropout),
|
||||||
|
nn.Linear(feature_dim, 128),
|
||||||
|
nn.ReLU(),
|
||||||
|
nn.Linear(128, num_classes_spine)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.hip_head = nn.Sequential(
|
||||||
|
nn.Dropout(dropout),
|
||||||
|
nn.Linear(feature_dim, 128),
|
||||||
|
nn.ReLU(),
|
||||||
|
nn.Linear(128, num_classes_hip)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.feature_dim = feature_dim
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor, region: str = 'spine') -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
Forward pass
|
||||||
|
|
||||||
|
Args:
|
||||||
|
x: [B, 3, H, W] Input image
|
||||||
|
region: 'spine' or 'hip'
|
||||||
|
"""
|
||||||
|
features = self.encoder(x)
|
||||||
|
|
||||||
|
if self.use_multi_task:
|
||||||
|
logits = self.head(features)
|
||||||
|
# 0-6: spine, 7-13: hip
|
||||||
|
if region == 'spine':
|
||||||
|
return logits[:, :self.num_classes_spine]
|
||||||
|
else:
|
||||||
|
return logits[:, self.num_classes_spine:]
|
||||||
|
else:
|
||||||
|
if region == 'spine':
|
||||||
|
return self.spine_head(features)
|
||||||
|
else:
|
||||||
|
return self.hip_head(features)
|
||||||
|
|
||||||
|
def predict(self,
|
||||||
|
image: torch.Tensor,
|
||||||
|
region: str = 'spine',
|
||||||
|
return_all: bool = False) -> Dict:
|
||||||
|
"""
|
||||||
|
Предсказание типа нарушения.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
image: Input tensor
|
||||||
|
region: 'spine' или 'hip'
|
||||||
|
return_all: вернуть все вероятности
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict с predicted_type, description, confidence, probabilities
|
||||||
|
"""
|
||||||
|
self.eval()
|
||||||
|
|
||||||
|
# Определение списка классов
|
||||||
|
violations = SPINE_VIOLATIONS if region == 'spine' else HIP_VIOLATIONS
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
logits = self.forward(image, region)
|
||||||
|
probs = torch.softmax(logits, dim=1)
|
||||||
|
conf, pred = probs.max(dim=1)
|
||||||
|
|
||||||
|
pred_idx = pred[0].item() if image.shape[0] == 1 else pred.item()
|
||||||
|
pred_type = violations[pred_idx]
|
||||||
|
|
||||||
|
result = {
|
||||||
|
'predicted_type': pred_type,
|
||||||
|
'description': VIOLATION_DESCRIPTIONS.get(pred_type, 'Неизвестно'),
|
||||||
|
'confidence': conf[0].item() if image.shape[0] == 1 else conf.item(),
|
||||||
|
'region': region,
|
||||||
|
'predicted_index': int(pred_idx)
|
||||||
|
}
|
||||||
|
|
||||||
|
if return_all:
|
||||||
|
result['probabilities'] = {
|
||||||
|
violations[i]: probs[0 if image.shape[0] == 1 else i, i].item()
|
||||||
|
for i in range(len(violations))
|
||||||
|
}
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
def predict_with_aggregated(self,
|
||||||
|
image: torch.Tensor,
|
||||||
|
region: str = 'spine') -> Dict:
|
||||||
|
"""
|
||||||
|
Предсказание с агрегированными результатами для батча.
|
||||||
|
"""
|
||||||
|
self.eval()
|
||||||
|
|
||||||
|
violations = SPINE_VIOLATIONS if region == 'spine' else HIP_VIOLATIONS
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
logits = self.forward(image, region)
|
||||||
|
probs = torch.softmax(logits, dim=1)
|
||||||
|
|
||||||
|
# Средние вероятности по батчу
|
||||||
|
avg_probs = probs.mean(dim=0)
|
||||||
|
conf, pred = avg_probs.max(dim=0)
|
||||||
|
|
||||||
|
return {
|
||||||
|
'predicted_type': violations[pred.item()],
|
||||||
|
'description': VIOLATION_DESCRIPTIONS.get(violations[pred.item()], 'Неизвестно'),
|
||||||
|
'confidence': conf.item(),
|
||||||
|
'predicted_index': pred.item(),
|
||||||
|
'probabilities': {
|
||||||
|
violations[i]: avg_probs[i].item()
|
||||||
|
for i in range(len(violations))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
def get_info(self) -> Dict:
|
||||||
|
return {
|
||||||
|
'name': 'ViolationTypeClassifier',
|
||||||
|
'backbone': self.backbone_name,
|
||||||
|
'num_classes_spine': self.num_classes_spine,
|
||||||
|
'num_classes_hip': self.num_classes_hip,
|
||||||
|
'spine_violations': SPINE_VIOLATIONS,
|
||||||
|
'hip_violations': HIP_VIOLATIONS,
|
||||||
|
'use_multi_task': self.use_multi_task
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def create_violation_classifier(
|
||||||
|
backbone: str = 'resnet18',
|
||||||
|
num_classes: int = 7,
|
||||||
|
pretrained: bool = True,
|
||||||
|
device: str = 'cpu'
|
||||||
|
) -> ViolationTypeClassifier:
|
||||||
|
"""Создание классификатора"""
|
||||||
|
model = ViolationTypeClassifier(
|
||||||
|
backbone=backbone,
|
||||||
|
num_classes_spine=num_classes,
|
||||||
|
num_classes_hip=num_classes,
|
||||||
|
pretrained=pretrained
|
||||||
|
)
|
||||||
|
model.to(device)
|
||||||
|
return model
|
||||||
|
|
@ -0,0 +1,192 @@
|
||||||
|
"""
|
||||||
|
Region Detector - определяет анатомическую область исследования
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torchvision.models as models
|
||||||
|
from typing import Dict, List, Tuple, Optional
|
||||||
|
import numpy as np
|
||||||
|
|
||||||
|
|
||||||
|
class RegionDetector(nn.Module):
|
||||||
|
"""
|
||||||
|
Определение анатомической области исследования.
|
||||||
|
Классы: spine, hip_left, hip_right, hip (unknown side)
|
||||||
|
"""
|
||||||
|
|
||||||
|
REGIONS = ['spine', 'hip_left', 'hip_right', 'hip']
|
||||||
|
|
||||||
|
def __init__(self,
|
||||||
|
backbone: str = 'efficientnet_b0',
|
||||||
|
num_classes: int = 4,
|
||||||
|
pretrained: bool = True,
|
||||||
|
dropout: float = 0.3):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.backbone_name = backbone
|
||||||
|
self.num_classes = num_classes
|
||||||
|
|
||||||
|
# Загрузка backbone
|
||||||
|
if backbone == 'efficientnet_b0':
|
||||||
|
self.backbone = models.efficientnet_b0(
|
||||||
|
weights='IMAGENET1K_V1' if pretrained else None
|
||||||
|
)
|
||||||
|
feature_dim = 1280
|
||||||
|
self.backbone.classifier = nn.Identity()
|
||||||
|
elif backbone == 'resnet18':
|
||||||
|
self.backbone = models.resnet18(
|
||||||
|
weights='IMAGENET1K_V1' if pretrained else None
|
||||||
|
)
|
||||||
|
feature_dim = 512
|
||||||
|
self.backbone.fc = nn.Identity()
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown backbone: {backbone}")
|
||||||
|
|
||||||
|
# Классификатор
|
||||||
|
self.classifier = nn.Sequential(
|
||||||
|
nn.Dropout(dropout),
|
||||||
|
nn.Linear(feature_dim, 256),
|
||||||
|
nn.ReLU(),
|
||||||
|
nn.Dropout(dropout),
|
||||||
|
nn.Linear(256, num_classes)
|
||||||
|
)
|
||||||
|
|
||||||
|
self.feature_dim = feature_dim
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
features = self.backbone(x)
|
||||||
|
return self.classifier(features)
|
||||||
|
|
||||||
|
def predict(self, x: torch.Tensor) -> Dict:
|
||||||
|
"""
|
||||||
|
Предсказание региона.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
x: Input tensor [B, 3, H, W]
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict с keys: predicted_region, probabilities, confidence
|
||||||
|
"""
|
||||||
|
self.eval()
|
||||||
|
with torch.no_grad():
|
||||||
|
logits = self.forward(x)
|
||||||
|
probs = torch.softmax(logits, dim=1)
|
||||||
|
conf, pred = probs.max(dim=1)
|
||||||
|
|
||||||
|
results = []
|
||||||
|
for i in range(x.shape[0]):
|
||||||
|
region_idx = pred[i].item()
|
||||||
|
results.append({
|
||||||
|
'predicted_region': self.REGIONS[region_idx],
|
||||||
|
'region_index': region_idx,
|
||||||
|
'confidence': conf[i].item(),
|
||||||
|
'probabilities': {
|
||||||
|
self.REGIONS[j]: probs[i, j].item()
|
||||||
|
for j in range(self.num_classes)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
return results[0] if len(results) == 1 else results
|
||||||
|
|
||||||
|
def get_info(self) -> Dict:
|
||||||
|
return {
|
||||||
|
'name': 'RegionDetector',
|
||||||
|
'backbone': self.backbone_name,
|
||||||
|
'num_classes': self.num_classes,
|
||||||
|
'regions': self.REGIONS,
|
||||||
|
'feature_dim': self.feature_dim
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class RegionDetectorWithFeatures(RegionDetector):
|
||||||
|
"""Region detector с извлечением признаков"""
|
||||||
|
|
||||||
|
def extract_features(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
return self.backbone(x)
|
||||||
|
|
||||||
|
|
||||||
|
def create_region_detector(
|
||||||
|
backbone: str = 'efficientnet_b0',
|
||||||
|
num_classes: int = 4,
|
||||||
|
pretrained: bool = True,
|
||||||
|
device: str = 'cpu'
|
||||||
|
) -> RegionDetector:
|
||||||
|
"""Создание модели"""
|
||||||
|
model = RegionDetector(
|
||||||
|
backbone=backbone,
|
||||||
|
num_classes=num_classes,
|
||||||
|
pretrained=pretrained
|
||||||
|
)
|
||||||
|
model.to(device)
|
||||||
|
return model
|
||||||
|
|
||||||
|
|
||||||
|
# Функция для определения региона из изображения (rule-based fallback)
|
||||||
|
def determine_region_from_image(image: np.ndarray) -> str:
|
||||||
|
"""
|
||||||
|
Определение региона из изображения без модели.
|
||||||
|
Использует форму яркой области.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
image: numpy array [H, W] или [H, W, 3]
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
str: 'spine', 'hip_left', 'hip_right', 'hip'
|
||||||
|
"""
|
||||||
|
if len(image.shape) == 3:
|
||||||
|
image = np.mean(image, axis=2)
|
||||||
|
|
||||||
|
# Нормализация
|
||||||
|
image = (image - image.min()) / (image.max() - image.min() + 1e-8)
|
||||||
|
|
||||||
|
h, w = image.shape
|
||||||
|
|
||||||
|
# Порог для яркой области
|
||||||
|
threshold = np.percentile(image, 95)
|
||||||
|
binary = image > threshold
|
||||||
|
|
||||||
|
if binary.sum() == 0:
|
||||||
|
return 'unknown'
|
||||||
|
|
||||||
|
# Aspect ratio яркой области
|
||||||
|
from scipy import ndimage
|
||||||
|
rows = np.any(binary, axis=1)
|
||||||
|
cols = np.any(binary, axis=0)
|
||||||
|
|
||||||
|
if not (rows.any() and cols.any()):
|
||||||
|
return 'unknown'
|
||||||
|
|
||||||
|
try:
|
||||||
|
rmin, rmax = np.where(rows)[0][[0, -1]]
|
||||||
|
cmin, cmax = np.where(cols)[0][[0, -1]]
|
||||||
|
|
||||||
|
bbox_h = rmax - rmin
|
||||||
|
bbox_w = cmax - cmin
|
||||||
|
bbox_aspect = bbox_h / (bbox_w + 1e-6)
|
||||||
|
|
||||||
|
# Симметрия
|
||||||
|
h_mid, w_mid = h // 2, w // 2
|
||||||
|
left_half = image[:, :w_mid]
|
||||||
|
right_half = np.fliplr(image[:, w_mid:])
|
||||||
|
min_w = min(left_half.shape[1], right_half.shape[1])
|
||||||
|
symmetry = 1 - np.abs(left_half[:, :min_w] - right_half[:, :min_w]).mean() / (image.std() + 1e-6)
|
||||||
|
|
||||||
|
# Left/right brightness ratio
|
||||||
|
left_bright = binary[:, :w//2].sum()
|
||||||
|
right_bright = binary[:, w//2:].sum()
|
||||||
|
left_right_ratio = left_bright / (right_bright + 1e-6)
|
||||||
|
|
||||||
|
# Классификация
|
||||||
|
if bbox_aspect < 1.5:
|
||||||
|
return 'spine'
|
||||||
|
elif bbox_aspect < 1.8:
|
||||||
|
return 'spine' if symmetry > 0.35 else 'hip'
|
||||||
|
else:
|
||||||
|
if left_right_ratio > 1.3:
|
||||||
|
return 'hip_right'
|
||||||
|
elif left_right_ratio < 0.7:
|
||||||
|
return 'hip_left'
|
||||||
|
else:
|
||||||
|
return 'hip'
|
||||||
|
except:
|
||||||
|
return 'unknown'
|
||||||
|
|
@ -0,0 +1,10 @@
|
||||||
|
"""
|
||||||
|
Segmentation models
|
||||||
|
"""
|
||||||
|
from .segmentator import DXASegmenter, DXASegmenterSimple, create_segmenter
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
'DXASegmenter',
|
||||||
|
'DXASegmenterSimple',
|
||||||
|
'create_segmenter'
|
||||||
|
]
|
||||||
|
|
@ -0,0 +1,254 @@
|
||||||
|
"""
|
||||||
|
U-Net сегментация для DXA изображений
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
import torchvision.models as models
|
||||||
|
from typing import Dict, Tuple, Optional, List
|
||||||
|
|
||||||
|
|
||||||
|
class DoubleConv(nn.Module):
|
||||||
|
"""Double convolution block"""
|
||||||
|
|
||||||
|
def __init__(self, in_channels, out_channels):
|
||||||
|
super().__init__()
|
||||||
|
self.conv = nn.Sequential(
|
||||||
|
nn.Conv2d(in_channels, out_channels, 3, padding=1),
|
||||||
|
nn.BatchNorm2d(out_channels),
|
||||||
|
nn.ReLU(inplace=True),
|
||||||
|
nn.Conv2d(out_channels, out_channels, 3, padding=1),
|
||||||
|
nn.BatchNorm2d(out_channels),
|
||||||
|
nn.ReLU(inplace=True)
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
return self.conv(x)
|
||||||
|
|
||||||
|
|
||||||
|
class Down(nn.Module):
|
||||||
|
"""Downsampling block"""
|
||||||
|
|
||||||
|
def __init__(self, in_channels, out_channels):
|
||||||
|
super().__init__()
|
||||||
|
self.maxpool_conv = nn.Sequential(
|
||||||
|
nn.MaxPool2d(2),
|
||||||
|
DoubleConv(in_channels, out_channels)
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
return self.maxpool_conv(x)
|
||||||
|
|
||||||
|
|
||||||
|
class Up(nn.Module):
|
||||||
|
"""Upsampling block"""
|
||||||
|
|
||||||
|
def __init__(self, in_channels, out_channels):
|
||||||
|
super().__init__()
|
||||||
|
self.up = nn.ConvTranspose2d(in_channels, in_channels // 2, 2, stride=2)
|
||||||
|
self.conv = DoubleConv(in_channels, out_channels)
|
||||||
|
|
||||||
|
def forward(self, x1, x2):
|
||||||
|
x1 = self.up(x1)
|
||||||
|
|
||||||
|
# Pad if sizes don't match
|
||||||
|
diffY = x2.size()[2] - x1.size()[2]
|
||||||
|
diffX = x2.size()[3] - x1.size()[3]
|
||||||
|
x1 = F.pad(x1, [diffX // 2, diffX - diffX // 2,
|
||||||
|
diffY // 2, diffY - diffY // 2])
|
||||||
|
|
||||||
|
x = torch.cat([x2, x1], dim=1)
|
||||||
|
return self.conv(x)
|
||||||
|
|
||||||
|
|
||||||
|
class OutConv(nn.Module):
|
||||||
|
"""Output convolution"""
|
||||||
|
|
||||||
|
def __init__(self, in_channels, out_channels):
|
||||||
|
super().__init__()
|
||||||
|
self.conv = nn.Conv2d(in_channels, out_channels, 1)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
return self.conv(x)
|
||||||
|
|
||||||
|
|
||||||
|
class DXASegmenter(nn.Module):
|
||||||
|
"""
|
||||||
|
U-Net сегментатор для DXA изображений.
|
||||||
|
Может работать для spine (позвонки) и hip (бедро).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self,
|
||||||
|
in_channels: int = 3,
|
||||||
|
out_classes: int = 2,
|
||||||
|
base_features: int = 64,
|
||||||
|
backbone: str = 'resnet34'):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.in_channels = in_channels
|
||||||
|
self.out_classes = out_classes
|
||||||
|
self.base_features = base_features
|
||||||
|
|
||||||
|
# Encoder (используем pretrained backbone)
|
||||||
|
if backbone == 'resnet34':
|
||||||
|
resnet = models.resnet34(weights='IMAGENET1K_V1')
|
||||||
|
self.encoder1 = nn.Sequential(resnet.conv1, resnet.bn1, resnet.relu)
|
||||||
|
self.encoder2 = nn.Sequential(resnet.maxpool, resnet.layer1)
|
||||||
|
self.encoder3 = resnet.layer2 # 256
|
||||||
|
self.encoder4 = resnet.layer3 # 512
|
||||||
|
self.encoder5 = resnet.layer4 # 1024
|
||||||
|
encoder_out = 1024
|
||||||
|
elif backbone == 'resnet18':
|
||||||
|
resnet = models.resnet18(weights='IMAGENET1K_V1')
|
||||||
|
self.encoder1 = nn.Sequential(resnet.conv1, resnet.bn1, resnet.relu)
|
||||||
|
self.encoder2 = nn.Sequential(resnet.maxpool, resnet.layer1)
|
||||||
|
self.encoder3 = resnet.layer2 # 128
|
||||||
|
self.encoder4 = resnet.layer3 # 256
|
||||||
|
self.encoder5 = resnet.layer4 # 512
|
||||||
|
encoder_out = 512
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown backbone: {backbone}")
|
||||||
|
|
||||||
|
# Decoder
|
||||||
|
self.up5 = Up(encoder_out, 512)
|
||||||
|
self.up4 = Up(512, 256)
|
||||||
|
self.up3 = Up(256, 128)
|
||||||
|
self.up2 = Up(128, 64)
|
||||||
|
|
||||||
|
self.out = OutConv(64, out_classes)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""
|
||||||
|
Forward pass
|
||||||
|
|
||||||
|
Args:
|
||||||
|
x: [B, 3, H, W] Input image
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
[B, out_classes, H, W] Segmentation logits
|
||||||
|
"""
|
||||||
|
# Encoder
|
||||||
|
x1 = self.encoder1(x) # 64
|
||||||
|
x2 = self.encoder2(x1) # 64
|
||||||
|
x3 = self.encoder3(x2) # 128/256
|
||||||
|
x4 = self.encoder4(x3) # 256/512
|
||||||
|
x5 = self.encoder5(x4) # 512/1024
|
||||||
|
|
||||||
|
# Decoder
|
||||||
|
d5 = self.up5(x5, x4)
|
||||||
|
d4 = self.up4(d5, x3)
|
||||||
|
d3 = self.up3(d4, x2)
|
||||||
|
d2 = self.up2(d3, x1)
|
||||||
|
|
||||||
|
return self.out(d2)
|
||||||
|
|
||||||
|
def predict(self, x: torch.Tensor) -> Dict:
|
||||||
|
"""
|
||||||
|
Предсказание сегментации.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict с mask, probabilities
|
||||||
|
"""
|
||||||
|
self.eval()
|
||||||
|
with torch.no_grad():
|
||||||
|
logits = self.forward(x)
|
||||||
|
probs = torch.softmax(logits, dim=1)
|
||||||
|
mask = probs.argmax(dim=1)
|
||||||
|
|
||||||
|
return {
|
||||||
|
'mask': mask, # [B, H, W]
|
||||||
|
'probabilities': probs, # [B, classes, H, W]
|
||||||
|
'logits': logits
|
||||||
|
}
|
||||||
|
|
||||||
|
def predict_mask(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
|
"""Получить только маску"""
|
||||||
|
result = self.predict(x)
|
||||||
|
return result['mask']
|
||||||
|
|
||||||
|
|
||||||
|
class DXASegmenterSimple(nn.Module):
|
||||||
|
"""
|
||||||
|
Упрощённая U-Net архитектура для сегментации.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self,
|
||||||
|
in_channels: int = 3,
|
||||||
|
out_classes: int = 2,
|
||||||
|
features: int = 64):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
self.in_channels = in_channels
|
||||||
|
self.out_classes = out_classes
|
||||||
|
|
||||||
|
# Encoder
|
||||||
|
self.inc = DoubleConv(in_channels, features)
|
||||||
|
self.down1 = Down(features, features * 2)
|
||||||
|
self.down2 = Down(features * 2, features * 4)
|
||||||
|
self.down3 = Down(features * 4, features * 8)
|
||||||
|
|
||||||
|
# Bottleneck
|
||||||
|
self.bottleneck = DoubleConv(features * 8, features * 16)
|
||||||
|
|
||||||
|
# Decoder
|
||||||
|
self.up3 = Up(features * 16, features * 8)
|
||||||
|
self.up2 = Up(features * 8, features * 4)
|
||||||
|
self.up1 = Up(features * 4, features * 2)
|
||||||
|
self.up0 = Up(features * 2, features)
|
||||||
|
|
||||||
|
self.out = OutConv(features, out_classes)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
x1 = self.inc(x)
|
||||||
|
x2 = self.down1(x1)
|
||||||
|
x3 = self.down2(x2)
|
||||||
|
x4 = self.down3(x3)
|
||||||
|
|
||||||
|
xb = self.bottleneck(x4)
|
||||||
|
|
||||||
|
d3 = self.up3(xb, x4)
|
||||||
|
d2 = self.up2(d3, x3)
|
||||||
|
d1 = self.up1(d2, x2)
|
||||||
|
d0 = self.up0(d1, x1)
|
||||||
|
|
||||||
|
return self.out(d0)
|
||||||
|
|
||||||
|
def predict(self, x: torch.Tensor) -> Dict:
|
||||||
|
self.eval()
|
||||||
|
with torch.no_grad():
|
||||||
|
logits = self.forward(x)
|
||||||
|
probs = torch.softmax(logits, dim=1)
|
||||||
|
mask = probs.argmax(dim=1)
|
||||||
|
|
||||||
|
return {
|
||||||
|
'mask': mask,
|
||||||
|
'probabilities': probs
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def create_segmenter(
|
||||||
|
architecture: str = 'unet_resnet34',
|
||||||
|
in_channels: int = 3,
|
||||||
|
out_classes: int = 2,
|
||||||
|
device: str = 'cpu'
|
||||||
|
) -> nn.Module:
|
||||||
|
"""
|
||||||
|
Создание сегментатора.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
architecture: 'unet_resnet34', 'unet_resnet18', 'unet_simple'
|
||||||
|
in_channels: количество входных каналов
|
||||||
|
out_classes: количество классов сегментации
|
||||||
|
device: устройство
|
||||||
|
"""
|
||||||
|
if architecture == 'unet_resnet34':
|
||||||
|
model = DXASegmenter(in_channels, out_classes, backbone='resnet34')
|
||||||
|
elif architecture == 'unet_resnet18':
|
||||||
|
model = DXASegmenter(in_channels, out_classes, backbone='resnet18')
|
||||||
|
elif architecture == 'unet_simple':
|
||||||
|
model = DXASegmenterSimple(in_channels, out_classes)
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unknown architecture: {architecture}")
|
||||||
|
|
||||||
|
model.to(device)
|
||||||
|
return model
|
||||||
|
|
@ -0,0 +1,10 @@
|
||||||
|
"""
|
||||||
|
Visualization models
|
||||||
|
"""
|
||||||
|
from .gradcam import GradCAMExtractor, SimpleGradCAM, visualize_attention
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
'GradCAMExtractor',
|
||||||
|
'SimpleGradCAM',
|
||||||
|
'visualize_attention'
|
||||||
|
]
|
||||||
|
|
@ -0,0 +1,249 @@
|
||||||
|
"""
|
||||||
|
Grad-CAM для визуализации attention модели
|
||||||
|
"""
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
import numpy as np
|
||||||
|
from typing import Dict, Optional, Tuple
|
||||||
|
import cv2
|
||||||
|
|
||||||
|
|
||||||
|
class GradCAMExtractor:
|
||||||
|
"""
|
||||||
|
Извлечение Grad-CAM heatmap для визуализации внимания модели.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self,
|
||||||
|
model: nn.Module,
|
||||||
|
target_layer: Optional[nn.Module] = None,
|
||||||
|
target_class: Optional[int] = None):
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
model: Модель для визуализации
|
||||||
|
target_layer: Слой для извлечения Grad-CAM
|
||||||
|
target_class: Класс для визуализации (None = pred class)
|
||||||
|
"""
|
||||||
|
self.model = model
|
||||||
|
self.target_layer = target_layer
|
||||||
|
self.target_class = target_class
|
||||||
|
self.gradients = None
|
||||||
|
self.activations = None
|
||||||
|
|
||||||
|
# Регистрация hooks
|
||||||
|
self._register_hooks()
|
||||||
|
|
||||||
|
def _register_hooks(self):
|
||||||
|
"""Регистрация forward и backward hooks"""
|
||||||
|
|
||||||
|
def forward_hook(module, input, output):
|
||||||
|
self.activations = output.detach()
|
||||||
|
|
||||||
|
def backward_hook(module, grad_input, grad_output):
|
||||||
|
self.gradients = grad_output[0].detach()
|
||||||
|
|
||||||
|
if self.target_layer is not None:
|
||||||
|
self.target_layer.register_forward_hook(forward_hook)
|
||||||
|
self.target_layer.register_full_backward_hook(backward_hook)
|
||||||
|
|
||||||
|
def set_target_layer(self, layer: nn.Module):
|
||||||
|
"""Установка целевого слоя"""
|
||||||
|
self.target_layer = layer
|
||||||
|
self._register_hooks()
|
||||||
|
|
||||||
|
def generate_cam(self,
|
||||||
|
input_tensor: torch.Tensor,
|
||||||
|
target_class: Optional[int] = None) -> np.ndarray:
|
||||||
|
"""
|
||||||
|
Генерация CAM heatmap.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
input_tensor: [1, C, H, W] Input image
|
||||||
|
target_class: Класс для визуализации
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Heatmap [H, W] normalized 0-1
|
||||||
|
"""
|
||||||
|
self.model.eval()
|
||||||
|
|
||||||
|
# Forward
|
||||||
|
output = self.model(input_tensor)
|
||||||
|
|
||||||
|
# Выбор класса
|
||||||
|
if target_class is None:
|
||||||
|
target_class = output.argmax(dim=1).item()
|
||||||
|
|
||||||
|
# Backward
|
||||||
|
self.model.zero_grad()
|
||||||
|
class_loss = output[0, target_class]
|
||||||
|
class_loss.backward()
|
||||||
|
|
||||||
|
# Вычисление CAM
|
||||||
|
if self.gradients is None or self.activations is None:
|
||||||
|
raise RuntimeError("Gradients or activations not captured")
|
||||||
|
|
||||||
|
# Global average pooling градиентов
|
||||||
|
pooled_gradients = torch.mean(self.gradients, dim=[0, 2, 3])
|
||||||
|
|
||||||
|
# Умножение активаций на градиенты
|
||||||
|
for i in range(self.activations.shape[1]):
|
||||||
|
self.activations[:, i, :, :] *= pooled_gradients[i]
|
||||||
|
|
||||||
|
# Average и ReLU
|
||||||
|
heatmap = torch.mean(self.activations, dim=1).squeeze()
|
||||||
|
heatmap = F.relu(heatmap)
|
||||||
|
|
||||||
|
# Normalization
|
||||||
|
if heatmap.max() > 0:
|
||||||
|
heatmap = heatmap / heatmap.max()
|
||||||
|
|
||||||
|
return heatmap.cpu().numpy()
|
||||||
|
|
||||||
|
def generate_cam_on_image(self,
|
||||||
|
input_tensor: torch.Tensor,
|
||||||
|
original_image: np.ndarray,
|
||||||
|
target_class: Optional[int] = None,
|
||||||
|
colormap: int = cv2.COLORMAP_JET) -> np.ndarray:
|
||||||
|
"""
|
||||||
|
Генерация CAM и наложение на оригинальное изображение.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
input_tensor: [1, C, H, W] Input tensor
|
||||||
|
original_image: [H, W, 3] RGB image (0-255)
|
||||||
|
target_class: Класс для визуализации
|
||||||
|
colormap: OpenCV colormap
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Overlay image [H, W, 3]
|
||||||
|
"""
|
||||||
|
# Get CAM
|
||||||
|
heatmap = self.generate_cam(input_tensor, target_class)
|
||||||
|
|
||||||
|
# Resize to original image
|
||||||
|
heatmap = cv2.resize(heatmap, (original_image.shape[1], original_image.shape[0]))
|
||||||
|
|
||||||
|
# Apply colormap
|
||||||
|
heatmap_colored = cv2.applyColorMap((heatmap * 255).astype(np.uint8), colormap)
|
||||||
|
heatmap_colored = cv2.cvtColor(heatmap_colored, cv2.COLOR_BGR2RGB)
|
||||||
|
|
||||||
|
# Overlay
|
||||||
|
overlay = cv2.addWeighted(original_image, 0.6, heatmap_colored, 0.4, 0)
|
||||||
|
|
||||||
|
return overlay
|
||||||
|
|
||||||
|
|
||||||
|
class SimpleGradCAM:
|
||||||
|
"""
|
||||||
|
Упрощённый Grad-CAM без hook - работает с любым классификатором.
|
||||||
|
"""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def generate(model: nn.Module,
|
||||||
|
input_tensor: torch.Tensor,
|
||||||
|
target_layer_name: str = 'encoder5') -> np.ndarray:
|
||||||
|
"""
|
||||||
|
Генерация heatmap.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model: Модель (должна иметь get_info() или быть классификатором)
|
||||||
|
input_tensor: [1, C, H, W]
|
||||||
|
target_layer_name: имя слоя для CAM
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Heatmap [H, W]
|
||||||
|
"""
|
||||||
|
model.eval()
|
||||||
|
|
||||||
|
# Для ResNet - последний слой
|
||||||
|
if hasattr(model, 'backbone'):
|
||||||
|
# Это наша модель
|
||||||
|
target_layer = model.backbone.layer4
|
||||||
|
else:
|
||||||
|
# Ищем последний conv слой
|
||||||
|
target_layer = None
|
||||||
|
for name, module in model.named_modules():
|
||||||
|
if isinstance(module, nn.Conv2d):
|
||||||
|
target_layer = module
|
||||||
|
|
||||||
|
if target_layer is None:
|
||||||
|
raise ValueError("Cannot find target layer")
|
||||||
|
|
||||||
|
extractor = GradCAMExtractor(model, target_layer)
|
||||||
|
|
||||||
|
# Forward + backward
|
||||||
|
output = model(input_tensor)
|
||||||
|
target_class = output.argmax(dim=1).item()
|
||||||
|
|
||||||
|
model.zero_grad()
|
||||||
|
output[0, target_class].backward()
|
||||||
|
|
||||||
|
# CAM
|
||||||
|
gradients = extractor.gradients
|
||||||
|
activations = extractor.activations
|
||||||
|
|
||||||
|
if gradients is None or activations is None:
|
||||||
|
raise RuntimeError("Failed to capture gradients")
|
||||||
|
|
||||||
|
pooled_gradients = torch.mean(gradients, dim=[0, 2, 3])
|
||||||
|
for i in range(activations.shape[1]):
|
||||||
|
activations[:, i, :, :] *= pooled_gradients[i]
|
||||||
|
|
||||||
|
heatmap = torch.mean(activations, dim=1).squeeze()
|
||||||
|
heatmap = F.relu(heatmap)
|
||||||
|
|
||||||
|
if heatmap.max() > 0:
|
||||||
|
heatmap = heatmap / heatmap.max()
|
||||||
|
|
||||||
|
return heatmap.cpu().numpy()
|
||||||
|
|
||||||
|
|
||||||
|
def visualize_attention(model: nn.Module,
|
||||||
|
image: torch.Tensor,
|
||||||
|
original_rgb: np.ndarray,
|
||||||
|
target_class: Optional[int] = None) -> Tuple[np.ndarray, np.ndarray]:
|
||||||
|
"""
|
||||||
|
Утилита для визуализации attention.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model: Классификатор
|
||||||
|
image: [1, 3, H, W] Tensor
|
||||||
|
original_rgb: [H, W, 3] RGB numpy
|
||||||
|
target_class: Опционально класс для визуализации
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
(heatmap, overlay) - heatmap и наложенное изображение
|
||||||
|
"""
|
||||||
|
# Try to use Grad-CAM
|
||||||
|
try:
|
||||||
|
# Find target layer
|
||||||
|
target_layer = None
|
||||||
|
for name, module in model.named_modules():
|
||||||
|
if hasattr(module, 'out_channels') and isinstance(module, nn.Conv2d):
|
||||||
|
target_layer = module
|
||||||
|
|
||||||
|
if target_layer is None:
|
||||||
|
# Use last conv layer of backbone
|
||||||
|
if hasattr(model, 'backbone'):
|
||||||
|
target_layer = model.backbone[-1]
|
||||||
|
|
||||||
|
if target_layer is not None:
|
||||||
|
extractor = GradCAMExtractor(model, target_layer)
|
||||||
|
heatmap = extractor.generate_cam(image, target_class)
|
||||||
|
|
||||||
|
# Resize
|
||||||
|
heatmap = cv2.resize(heatmap, (original_rgb.shape[1], original_rgb.shape[0]))
|
||||||
|
|
||||||
|
# Overlay
|
||||||
|
heatmap_colored = cv2.applyColorMap(
|
||||||
|
(heatmap * 255).astype(np.uint8),
|
||||||
|
cv2.COLORMAP_JET
|
||||||
|
)
|
||||||
|
heatmap_colored = cv2.cvtColor(heatmap_colored, cv2.COLOR_BGR2RGB)
|
||||||
|
overlay = cv2.addWeighted(original_rgb, 0.6, heatmap_colored, 0.4, 0)
|
||||||
|
|
||||||
|
return heatmap, overlay
|
||||||
|
except Exception as e:
|
||||||
|
print(f"Grad-CAM failed: {e}")
|
||||||
|
|
||||||
|
# Fallback: return original image
|
||||||
|
return np.zeros((original_rgb.shape[0], original_rgb.shape[1])), original_rgb
|
||||||
|
|
@ -0,0 +1,16 @@
|
||||||
|
"""
|
||||||
|
Pipeline components
|
||||||
|
"""
|
||||||
|
from .pipeline import (
|
||||||
|
QualityAssessmentPipeline,
|
||||||
|
PipelineConfig,
|
||||||
|
QualityResult,
|
||||||
|
create_pipeline
|
||||||
|
)
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
'QualityAssessmentPipeline',
|
||||||
|
'PipelineConfig',
|
||||||
|
'QualityResult',
|
||||||
|
'create_pipeline'
|
||||||
|
]
|
||||||
|
|
@ -0,0 +1,289 @@
|
||||||
|
"""
|
||||||
|
Основной пайплайн оценки качества 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)
|
||||||
|
|
@ -0,0 +1,281 @@
|
||||||
|
"""
|
||||||
|
Training script for DXA Quality Models
|
||||||
|
"""
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import argparse
|
||||||
|
from pathlib import Path
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch.utils.data import Dataset, DataLoader, Subset
|
||||||
|
from PIL import Image
|
||||||
|
import pydicom
|
||||||
|
from tqdm import tqdm
|
||||||
|
import random
|
||||||
|
|
||||||
|
# Add project root to path
|
||||||
|
project_root = Path(__file__).parent.parent.parent
|
||||||
|
sys.path.insert(0, str(project_root))
|
||||||
|
os.chdir(project_root)
|
||||||
|
|
||||||
|
|
||||||
|
class DXATrainingDataset(Dataset):
|
||||||
|
"""Dataset for training DXA models"""
|
||||||
|
|
||||||
|
def __init__(self, data_root, annotation_path, region=None, input_size=224):
|
||||||
|
self.data_root = Path(data_root)
|
||||||
|
self.annotation_path = annotation_path
|
||||||
|
self.region = region
|
||||||
|
self.input_size = input_size
|
||||||
|
|
||||||
|
self.annotation = self._load_annotation()
|
||||||
|
self.samples = self._build_samples()
|
||||||
|
|
||||||
|
print(f"Loaded {len(self.samples)} samples for region={region}")
|
||||||
|
|
||||||
|
def _load_annotation(self):
|
||||||
|
df = pd.read_excel(self.annotation_path, header=None)
|
||||||
|
data = df.iloc[2:].copy()
|
||||||
|
data.columns = range(len(df.columns))
|
||||||
|
data = data.rename(columns={
|
||||||
|
0: 'id',
|
||||||
|
1: 'study_uid',
|
||||||
|
9: 'spine_total',
|
||||||
|
10: 'hip_right_total',
|
||||||
|
11: 'hip_left_total'
|
||||||
|
})
|
||||||
|
data = data.dropna(subset=['study_uid'])
|
||||||
|
return data
|
||||||
|
|
||||||
|
def _build_samples(self):
|
||||||
|
samples = []
|
||||||
|
|
||||||
|
for _, row in self.annotation.iterrows():
|
||||||
|
study_uid = str(row['study_uid']).strip()
|
||||||
|
|
||||||
|
# Find study folder - it's inside Исследования subfolder
|
||||||
|
study_path = self.data_root / 'Исследования' / study_uid
|
||||||
|
|
||||||
|
if not study_path.exists():
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Find DICOM files
|
||||||
|
dcm_files = sorted(study_path.rglob('*.dcm'))
|
||||||
|
|
||||||
|
for dcm_file in dcm_files:
|
||||||
|
fname = dcm_file.name.lower()
|
||||||
|
|
||||||
|
# Determine region from filename
|
||||||
|
if 'spine' in fname:
|
||||||
|
sample_region = 'spine'
|
||||||
|
target_col = 'spine_total'
|
||||||
|
elif 'l_hip' in fname or 'left' in fname:
|
||||||
|
sample_region = 'hip_left'
|
||||||
|
target_col = 'hip_left_total'
|
||||||
|
elif 'r_hip' in fname or 'right' in fname:
|
||||||
|
sample_region = 'hip_right'
|
||||||
|
target_col = 'hip_right_total'
|
||||||
|
else:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Filter by region if specified
|
||||||
|
if self.region and sample_region != self.region:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Get target from the same row
|
||||||
|
if target_col in row and pd.notna(row[target_col]):
|
||||||
|
target = int(row[target_col])
|
||||||
|
else:
|
||||||
|
# Use spine total as fallback
|
||||||
|
if pd.notna(row.get('spine_total', None)):
|
||||||
|
target = int(row['spine_total'])
|
||||||
|
else:
|
||||||
|
continue
|
||||||
|
|
||||||
|
samples.append({
|
||||||
|
'path': str(dcm_file),
|
||||||
|
'region': sample_region,
|
||||||
|
'target': target
|
||||||
|
})
|
||||||
|
|
||||||
|
return samples
|
||||||
|
|
||||||
|
def __len__(self):
|
||||||
|
return len(self.samples)
|
||||||
|
|
||||||
|
def __getitem__(self, idx):
|
||||||
|
sample = self.samples[idx]
|
||||||
|
|
||||||
|
# Load DICOM
|
||||||
|
ds = pydicom.dcmread(sample['path'])
|
||||||
|
img = ds.pixel_array.astype(np.float32)
|
||||||
|
|
||||||
|
# Normalize
|
||||||
|
img = (img - img.min()) / (img.max() - img.min() + 1e-8)
|
||||||
|
|
||||||
|
# Convert to RGB
|
||||||
|
img = np.stack([img] * 3, axis=2)
|
||||||
|
|
||||||
|
# Resize
|
||||||
|
img_pil = Image.fromarray((img * 255).astype(np.uint8))
|
||||||
|
img_pil = img_pil.resize((self.input_size, self.input_size), Image.BILINEAR)
|
||||||
|
img = np.array(img_pil).astype(np.float32) / 255.0
|
||||||
|
|
||||||
|
# CHW
|
||||||
|
img = img.transpose(2, 0, 1)
|
||||||
|
img = torch.from_numpy(img).float()
|
||||||
|
|
||||||
|
return img, torch.tensor(sample['target'], dtype=torch.long)
|
||||||
|
|
||||||
|
|
||||||
|
def train_model(model, train_loader, val_loader, epochs, lr, device, save_path):
|
||||||
|
"""Train model"""
|
||||||
|
model = model.to(device)
|
||||||
|
criterion = nn.CrossEntropyLoss()
|
||||||
|
optimizer = torch.optim.AdamW(model.parameters(), lr=lr, weight_decay=1e-4)
|
||||||
|
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=2)
|
||||||
|
|
||||||
|
best_val_acc = 0
|
||||||
|
|
||||||
|
for epoch in range(epochs):
|
||||||
|
# Train
|
||||||
|
model.train()
|
||||||
|
train_loss = 0
|
||||||
|
train_correct = 0
|
||||||
|
train_total = 0
|
||||||
|
|
||||||
|
pbar = tqdm(train_loader, desc=f"Epoch {epoch+1}/{epochs}")
|
||||||
|
for images, labels in pbar:
|
||||||
|
images = images.to(device)
|
||||||
|
labels = labels.to(device)
|
||||||
|
|
||||||
|
optimizer.zero_grad()
|
||||||
|
outputs = model(images)
|
||||||
|
loss = criterion(outputs, labels)
|
||||||
|
loss.backward()
|
||||||
|
optimizer.step()
|
||||||
|
|
||||||
|
train_loss += loss.item() * images.size(0)
|
||||||
|
_, predicted = outputs.max(1)
|
||||||
|
train_correct += predicted.eq(labels).sum().item()
|
||||||
|
train_total += labels.size(0)
|
||||||
|
|
||||||
|
pbar.set_postfix({'loss': train_loss/train_total, 'acc': train_correct/train_total})
|
||||||
|
|
||||||
|
train_acc = train_correct / train_total
|
||||||
|
|
||||||
|
# Validate
|
||||||
|
model.eval()
|
||||||
|
val_loss = 0
|
||||||
|
val_correct = 0
|
||||||
|
val_total = 0
|
||||||
|
|
||||||
|
with torch.no_grad():
|
||||||
|
for images, labels in val_loader:
|
||||||
|
images = images.to(device)
|
||||||
|
labels = labels.to(device)
|
||||||
|
|
||||||
|
outputs = model(images)
|
||||||
|
loss = criterion(outputs, labels)
|
||||||
|
|
||||||
|
val_loss += loss.item() * images.size(0)
|
||||||
|
_, predicted = outputs.max(1)
|
||||||
|
val_correct += predicted.eq(labels).sum().item()
|
||||||
|
val_total += labels.size(0)
|
||||||
|
|
||||||
|
val_acc = val_correct / val_total
|
||||||
|
|
||||||
|
scheduler.step(val_loss)
|
||||||
|
|
||||||
|
print(f"Epoch {epoch+1}: Train Acc={train_acc:.3f}, Val Acc={val_acc:.3f}, Val Loss={val_loss/val_total:.4f}")
|
||||||
|
|
||||||
|
if val_acc > best_val_acc:
|
||||||
|
best_val_acc = val_acc
|
||||||
|
torch.save(model.state_dict(), save_path)
|
||||||
|
print(f" -> Saved best model to {save_path}")
|
||||||
|
|
||||||
|
return best_val_acc
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
parser = argparse.ArgumentParser()
|
||||||
|
parser.add_argument('--data-root', default='dataset_hack/НД_для_обучения')
|
||||||
|
parser.add_argument('--annotation', default='dataset_hack/НД_для_обучения/разметка.xlsx')
|
||||||
|
parser.add_argument('--model-type', default='quality', choices=['quality', 'region', 'violation'])
|
||||||
|
parser.add_argument('--epochs', type=int, default=15)
|
||||||
|
parser.add_argument('--batch-size', type=int, default=8)
|
||||||
|
parser.add_argument('--lr', type=float, default=1e-4)
|
||||||
|
parser.add_argument('--input-size', type=int, default=224)
|
||||||
|
parser.add_argument('--output', default='models/')
|
||||||
|
parser.add_argument('--region', default=None, help='spine, hip_left, hip_right, or None for all')
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
# Device
|
||||||
|
device = torch.device('cuda' if torch.cuda.is_available() else 'mps' if torch.backends.mps.is_available() else 'cpu')
|
||||||
|
print(f"Using device: {device}")
|
||||||
|
|
||||||
|
# Create output dir
|
||||||
|
os.makedirs(args.output, exist_ok=True)
|
||||||
|
|
||||||
|
# Load dataset
|
||||||
|
dataset = DXATrainingDataset(
|
||||||
|
data_root=args.data_root,
|
||||||
|
annotation_path=args.annotation,
|
||||||
|
region=args.region,
|
||||||
|
input_size=args.input_size
|
||||||
|
)
|
||||||
|
|
||||||
|
if len(dataset) == 0:
|
||||||
|
print("ERROR: No samples found!")
|
||||||
|
return
|
||||||
|
|
||||||
|
# Split
|
||||||
|
n = len(dataset)
|
||||||
|
n_train = int(n * 0.8)
|
||||||
|
n_val = n - n_train
|
||||||
|
|
||||||
|
# Random split
|
||||||
|
indices = list(range(n))
|
||||||
|
random.shuffle(indices)
|
||||||
|
train_indices = indices[:n_train]
|
||||||
|
val_indices = indices[n_train:]
|
||||||
|
|
||||||
|
train_dataset = Subset(dataset, train_indices)
|
||||||
|
val_dataset = Subset(dataset, val_indices)
|
||||||
|
|
||||||
|
train_loader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True, num_workers=0)
|
||||||
|
val_loader = DataLoader(val_dataset, batch_size=args.batch_size, shuffle=False, num_workers=0)
|
||||||
|
|
||||||
|
print(f"Train: {len(train_dataset)}, Val: {len(val_dataset)}")
|
||||||
|
|
||||||
|
# Class distribution
|
||||||
|
targets = [dataset.samples[i]['target'] for i in train_indices]
|
||||||
|
n_pos = sum(targets)
|
||||||
|
n_neg = len(targets) - n_pos
|
||||||
|
print(f"Class distribution - Positive: {n_pos}, Negative: {n_neg}")
|
||||||
|
|
||||||
|
# Create model
|
||||||
|
if args.model_type == 'quality':
|
||||||
|
from src.models.classification import QualityClassifier
|
||||||
|
model = QualityClassifier(backbone='resnet18', pretrained=True)
|
||||||
|
save_path = f"{args.output}/quality_classifier.pth"
|
||||||
|
elif args.model_type == 'region':
|
||||||
|
from src.models import RegionDetector
|
||||||
|
model = RegionDetector(backbone='resnet18', pretrained=True)
|
||||||
|
save_path = f"{args.output}/region_detector.pth"
|
||||||
|
else:
|
||||||
|
from src.models.classification import ViolationTypeClassifier
|
||||||
|
model = ViolationTypeClassifier(backbone='resnet18', pretrained=True)
|
||||||
|
save_path = f"{args.output}/violation_classifier.pth"
|
||||||
|
|
||||||
|
# Train
|
||||||
|
print(f"Training {args.model_type} model...")
|
||||||
|
best_acc = train_model(model, train_loader, val_loader, args.epochs, args.lr, device, save_path)
|
||||||
|
print(f"Best validation accuracy: {best_acc:.3f}")
|
||||||
|
print(f"Model saved to: {save_path}")
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
main()
|
||||||
Loading…
Reference in New Issue