238 lines
8.2 KiB
Python
238 lines
8.2 KiB
Python
"""
|
||
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
|