277 lines
12 KiB
Python
277 lines
12 KiB
Python
# src/models/bone_condition_classifier.py
|
||
from typing import Dict, Any, Union, Optional, List
|
||
import torch
|
||
import torch.nn as nn
|
||
import torchvision.models as models
|
||
from PIL import Image
|
||
import numpy as np
|
||
from pathlib import Path
|
||
|
||
from src.classifiers.base import BaseClassifier
|
||
|
||
# Маппинг состояний костей
|
||
BONE_CONDITIONS = {
|
||
0: "Normal",
|
||
1: "Osteopenia",
|
||
2: "Osteoporosis",
|
||
3: "Severe Osteoporosis",
|
||
4: "Fracture",
|
||
5: "Degenerative Changes",
|
||
6: "Arthritis",
|
||
7: "Tumor",
|
||
8: "Infection",
|
||
9: "Post-surgical Changes",
|
||
10: "Congenital Anomaly"
|
||
}
|
||
|
||
SEVERITY_MAP = {
|
||
"Normal": "None",
|
||
"Osteopenia": "Mild",
|
||
"Osteoporosis": "Moderate",
|
||
"Severe Osteoporosis": "Severe",
|
||
"Fracture": "Acute",
|
||
"Degenerative Changes": "Chronic",
|
||
"Arthritis": "Chronic",
|
||
"Tumor": "Severe",
|
||
"Infection": "Severe",
|
||
"Post-surgical Changes": "Mild",
|
||
"Congenital Anomaly": "Moderate"
|
||
}
|
||
|
||
|
||
class BoneConditionClassifier(BaseClassifier):
|
||
"""
|
||
Классификатор состояния костной ткани
|
||
Используется для медицинской диагностики на основе денситометрических изображений
|
||
"""
|
||
|
||
def __init__(
|
||
self,
|
||
model_path: Optional[str] = None,
|
||
device: str = 'cpu',
|
||
num_classes: int = 11, # 10 состояний + 1 норма
|
||
input_size: int = 224
|
||
):
|
||
"""
|
||
Args:
|
||
model_path: путь к обученной модели (опционально)
|
||
device: устройство для инференса ('cpu', 'mps', 'cuda')
|
||
num_classes: количество классов состояний
|
||
input_size: размер входного изображения
|
||
"""
|
||
self.device = device
|
||
self.input_size = input_size
|
||
self.num_classes = num_classes
|
||
|
||
# Используем предобученный ResNet50 как бэкбон
|
||
self.model = models.resnet50(pretrained=True)
|
||
|
||
# Заменяем последний слой на количество состояний
|
||
num_features = self.model.fc.in_features
|
||
self.model.fc = nn.Sequential(
|
||
nn.Dropout(0.5),
|
||
nn.Linear(num_features, 256),
|
||
nn.ReLU(),
|
||
nn.Dropout(0.3),
|
||
nn.Linear(256, num_classes)
|
||
)
|
||
|
||
# Загружаем веса если есть
|
||
self.loaded = False
|
||
if model_path and Path(model_path).exists():
|
||
try:
|
||
self.model.load_state_dict(torch.load(model_path, map_location=device))
|
||
self.loaded = True
|
||
print(f"✅ Загружен классификатор костей из {model_path}")
|
||
except Exception as e:
|
||
print(f"⚠️ Ошибка загрузки модели из {model_path}: {e}")
|
||
print(" Используем неподготовленную модель (будет работать плохо)")
|
||
else:
|
||
print("⚠️ Классификатор костей не найден")
|
||
print(f" Ожидается: {model_path}")
|
||
print(" Обучите модель на медицинских данных перед использованием")
|
||
|
||
self.model.to(device)
|
||
self.model.eval()
|
||
|
||
# Нормализация для медицинских изображений (можно адаптировать под DICOM)
|
||
self.mean = np.array([0.485, 0.456, 0.406])
|
||
self.std = np.array([0.229, 0.224, 0.225])
|
||
|
||
print(f"📊 Модель на устройстве: {device}")
|
||
print(f"📋 Количество классов: {num_classes}")
|
||
|
||
def preprocess(self, image: Union[np.ndarray, Image.Image]) -> torch.Tensor:
|
||
"""
|
||
Предобработка изображения для модели
|
||
|
||
Args:
|
||
image: входное изображение
|
||
|
||
Returns:
|
||
torch.Tensor: подготовленный тензор
|
||
"""
|
||
# Конвертация в PIL Image
|
||
if isinstance(image, np.ndarray):
|
||
image = Image.fromarray(image)
|
||
|
||
# Ресайз
|
||
image = image.resize((self.input_size, self.input_size))
|
||
|
||
# Конвертация в массив и нормализация
|
||
image_array = np.array(image, dtype=np.float32) / 255.0
|
||
|
||
# Если изображение grayscale, конвертируем в RGB
|
||
if len(image_array.shape) == 2:
|
||
image_array = np.stack([image_array] * 3, axis=2)
|
||
elif image_array.shape[2] == 1:
|
||
image_array = np.concatenate([image_array] * 3, axis=2)
|
||
|
||
# Изменение порядка каналов HWC -> CHW
|
||
image_array = image_array.transpose(2, 0, 1)
|
||
|
||
# Нормализация
|
||
for i in range(3):
|
||
image_array[i] = (image_array[i] - self.mean[i]) / self.std[i]
|
||
|
||
return torch.FloatTensor(image_array).unsqueeze(0).to(self.device)
|
||
|
||
def predict(self, image: Union[np.ndarray, Image.Image]) -> Dict[str, Any]:
|
||
"""
|
||
Предсказание состояния костной ткани
|
||
|
||
Args:
|
||
image: изображение в формате numpy array или PIL Image
|
||
|
||
Returns:
|
||
Словарь с результатами:
|
||
- condition_id: ID состояния
|
||
- condition_name: название состояния
|
||
- severity: серьезность состояния
|
||
- confidence: уверенность (0-1)
|
||
- all_probs: вероятности всех классов
|
||
- loaded: загружена ли модель
|
||
"""
|
||
# Проверка загрузки модели
|
||
if not self.loaded:
|
||
return {
|
||
"condition_id": -1,
|
||
"condition_name": "Unknown",
|
||
"severity": "Unknown",
|
||
"confidence": 0.0,
|
||
"loaded": False,
|
||
"error": "Model not loaded"
|
||
}
|
||
|
||
try:
|
||
# Предобработка
|
||
image_tensor = self.preprocess(image)
|
||
|
||
# Инференс
|
||
with torch.no_grad():
|
||
outputs = self.model(image_tensor)
|
||
probabilities = torch.softmax(outputs, dim=1)
|
||
confidence, predicted = torch.max(probabilities, 1)
|
||
|
||
# Формируем результат
|
||
condition_id = predicted.item()
|
||
confidence_score = confidence.item()
|
||
all_probs = probabilities.cpu().numpy().tolist()[0]
|
||
|
||
condition_name = BONE_CONDITIONS.get(condition_id, "Unknown")
|
||
severity = SEVERITY_MAP.get(condition_name, "Unknown")
|
||
|
||
# Дополнительные метрики для медицинского контекста
|
||
is_abnormal = condition_id != 0
|
||
risk_level = "Low"
|
||
if condition_id in [2, 3, 7, 8]: # Остеопороз, опухоль, инфекция
|
||
risk_level = "High"
|
||
elif condition_id in [1, 4, 5, 6]: # Остеопения, перелом, дегенерация
|
||
risk_level = "Medium"
|
||
|
||
return {
|
||
"condition_id": condition_id,
|
||
"condition_name": condition_name,
|
||
"severity": severity,
|
||
"confidence": confidence_score,
|
||
"all_probs": all_probs,
|
||
"is_abnormal": is_abnormal,
|
||
"risk_level": risk_level,
|
||
"loaded": True,
|
||
"error": None
|
||
}
|
||
|
||
except Exception as e:
|
||
return {
|
||
"condition_id": -1,
|
||
"condition_name": "Error",
|
||
"severity": "Unknown",
|
||
"confidence": 0.0,
|
||
"loaded": self.loaded,
|
||
"error": str(e)
|
||
}
|
||
|
||
def classify(self, image: Union[np.ndarray, Image.Image]) -> Dict[str, Any]:
|
||
"""Алиас для predict (для совместимости с BaseClassifier)"""
|
||
return self.predict(image)
|
||
|
||
def predict_batch(self, images: List[Union[np.ndarray, Image.Image]]) -> List[Dict[str, Any]]:
|
||
"""
|
||
Предсказание для батча изображений
|
||
|
||
Args:
|
||
images: список изображений
|
||
|
||
Returns:
|
||
Список результатов для каждого изображения
|
||
"""
|
||
return [self.predict(img) for img in images]
|
||
|
||
def get_info(self) -> Dict[str, Any]:
|
||
"""Информация о классификаторе"""
|
||
return {
|
||
"name": self.__class__.__name__,
|
||
"type": "classification",
|
||
"target": "bone_conditions",
|
||
"num_classes": self.num_classes,
|
||
"loaded": self.loaded,
|
||
"device": self.device,
|
||
"model_architecture": "ResNet50 + Custom Head",
|
||
"pretrained": True,
|
||
"input_size": self.input_size,
|
||
"conditions": list(BONE_CONDITIONS.values()),
|
||
"severity_map": SEVERITY_MAP
|
||
}
|
||
|
||
def get_recommendations(self, condition_id: int) -> List[str]:
|
||
"""
|
||
Получение рекомендаций на основе состояния
|
||
|
||
Args:
|
||
condition_id: ID состояния
|
||
|
||
Returns:
|
||
Список рекомендаций
|
||
"""
|
||
recommendations = {
|
||
0: ["✅ Состояние в норме", "Продолжайте мониторинг"],
|
||
1: ["⚠️ Начальные изменения плотности", "Рекомендуется контроль через 6 месяцев",
|
||
"Увеличьте потребление кальция"],
|
||
2: ["🔴 Умеренное снижение плотности", "Рекомендуется контроль через 3 месяца", "Консультация эндокринолога",
|
||
"Препараты кальция и витамин D"],
|
||
3: ["🚨 Критическое снижение плотности", "Немедленная консультация специалиста", "Интенсивная терапия",
|
||
"Мониторинг переломов"],
|
||
4: ["🦴 Обнаружен перелом", "Иммобилизация", "Консультация травматолога", "Контрольная рентгенография"],
|
||
5: ["⚙️ Дегенеративные изменения", "Физиотерапия", "Противовоспалительная терапия",
|
||
"Контроль через 6 месяцев"],
|
||
6: ["🔄 Артрит", "Противовоспалительная терапия", "Физиотерапия", "Консультация ревматолога"],
|
||
7: ["🧬 Подозрение на опухоль", "Срочная консультация онколога", "МРТ/КТ исследование", "Биопсия"],
|
||
8: ["🦠 Подозрение на инфекцию", "Антибактериальная терапия", "Консультация инфекциониста",
|
||
"Контрольный анализ"],
|
||
9: ["🔧 Послеоперационные изменения", "Контроль через 3 месяца", "Физиотерапия", "Наблюдение хирурга"],
|
||
10: ["🧬 Врожденная аномалия", "Консультация генетика", "Индивидуальный план лечения", "Мониторинг развития"]
|
||
}
|
||
return recommendations.get(condition_id, ["Обратитесь к специалисту"])
|
||
|
||
|