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