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