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