bone_2026/src/models/classification/violation_classifier.py

238 lines
8.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

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

"""
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