bone_2026/src/models/region_detector.py

193 lines
6.0 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.

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