bone_2026/src/classifiers/bone_condition_classifier.py

277 lines
12 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.

# 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, ["Обратитесь к специалисту"])