193 lines
6.0 KiB
Python
193 lines
6.0 KiB
Python
"""
|
||
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'
|