Enterwar99/mammography-api-server
0
1from fastapi import FastAPI, File, UploadFile, HTTPException2from fastapi.responses import JSONResponse3import torch4import torchvision.models as models5import torchvision.transforms as transforms6from PIL import Image7import torch.nn as nn8import io9import numpy as np10import os11from typing import List, Dict, Any, Optional12import logging13import cv214import base6415 16from pytorch_grad_cam import GradCAMPlusPlus17from pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget18from huggingface_hub import hf_hub_download19from pydantic import BaseModel20 21# --- Konfiguracja Logowania ---22logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')23logger = logging.getLogger(__name__)24 25# --- Konfiguracja ---26HF_MODEL_REPO_ID = "Enterwar99/MODEL_MAMMOGRAFII"27MODEL_FILENAME = "best_model.pth"28 29DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")30IMAGENET_MEAN = [0.485, 0.456, 0.406]31IMAGENET_STD = [0.229, 0.224, 0.225]32 33# Globalne zmienne dla modelu i transformacji34model_instance = None35transform_pipeline = None36 37interpretations_dict = {38 1: "Wynik negatywny - brak zmian nowotworowych",39 2: "Zmiana łagodna",40 3: "Prawdopodobnie zmiana łagodna - zalecana kontrola",41 4: "Podejrzenie zmiany złośliwej - zalecana biopsja",42 5: "Wysoka podejrzliwość złośliwości - wymagana biopsja"43}44 45# --- Inicjalizacja modelu ---46def initialize_model():47 global model_instance, transform_pipeline48 if model_instance is not None:49 return50 51 logger.info("Rozpoczynanie inicjalizacji modelu...")52 try:53 hf_auth_token = os.environ.get("HF_TOKEN_MODEL_READ")54 model_pt_path = hf_hub_download(repo_id=HF_MODEL_REPO_ID, filename=MODEL_FILENAME, token=hf_auth_token)55 logger.info(f"Plik modelu pomyślnie pobrany do: {model_pt_path}")56 except Exception as e:57 logger.error(f"Błąd podczas pobierania modelu z Hugging Face Hub: {e}", exc_info=True)58 raise RuntimeError(f"Nie można pobrać modelu: {e}")59 60 model_arch = models.resnet18(weights=None)61 num_feats = model_arch.fc.in_features62 model_arch.fc = nn.Sequential(nn.Dropout(0.5), nn.Linear(num_feats, 5))63 64 model_arch.load_state_dict(torch.load(model_pt_path, map_location=DEVICE))65 model_arch.to(DEVICE)66 model_arch.eval()67 model_instance = model_arch68 69 transform_pipeline = transforms.Compose([70 transforms.Resize((224, 224)),71 transforms.ToTensor(),72 transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD)73 ])74 logger.info(f"Model BI-RADS classifier initialized successfully on device: {DEVICE}")75 76# --- Funkcja do predykcji z kwantyfikacją niepewności (MC Dropout) ---77def predict_with_mc_dropout(current_model_instance, batch_tensor_on_device, mc_dropout_samples: int, uncertainty_threshold_std: float):78 logger.info(f"Performing MC Dropout on a batch of size {batch_tensor_on_device.shape[0]} with {mc_dropout_samples} samples.")79 80 original_mode_is_training = current_model_instance.training81 current_model_instance.train()82 83 batch_size = batch_tensor_on_device.shape[0]84 num_classes = 585 86 all_probs_batch = np.zeros((batch_size, mc_dropout_samples, num_classes))87 88 with torch.no_grad():89 for i in range(mc_dropout_samples):90 output = current_model_instance(batch_tensor_on_device)91 probs_tensor = torch.nn.functional.softmax(output, dim=1)92 all_probs_batch[:, i, :] = probs_tensor.cpu().numpy()93 94 if not original_mode_is_training:95 current_model_instance.eval()96 97 mean_probabilities_batch = np.mean(all_probs_batch, axis=1)98 std_dev_probabilities_batch = np.std(all_probs_batch, axis=1)99 100 results = []101 for i in range(batch_size):102 mean_probabilities = mean_probabilities_batch[i]103 std_dev_probabilities = std_dev_probabilities_batch[i]104 105 predicted_class_index = np.argmax(mean_probabilities)106 confidence_in_predicted_class = float(np.max(all_probs_batch[i, :, predicted_class_index]))107 uncertainty_metric = np.mean(std_dev_probabilities)108 is_uncertain = uncertainty_metric > uncertainty_threshold_std109 110 logger.info(f"MC Dropout Results for image {i}: Predicted Index: {int(predicted_class_index)}, Confidence (MaxProb): {confidence_in_predicted_class:.4f}, Uncertainty (avg_std): {uncertainty_metric:.4f}, Is Uncertain: {is_uncertain}")111 112 birads_category_if_confident = int(predicted_class_index) + 1113 114 if is_uncertain:115 result = {116 "birads": None, "confidence": None,117 "interpretation": f"Model jest niepewny co do tego obrazu (niepewność: {uncertainty_metric:.4f}). Sprawdź jakość i typ obrazu.",118 "class_probabilities": {str(j + 1): float(mean_probabilities[j]) for j in range(len(mean_probabilities))},119 "grad_cam_image_base64": None, "error": "High prediction uncertainty",120 "details": f"Uncertainty metric ({uncertainty_metric:.4f}) przekroczyła próg ({uncertainty_threshold_std})."121 }122 else:123 result = {124 "birads": birads_category_if_confident,125 "confidence": confidence_in_predicted_class,126 "interpretation": interpretations_dict.get(birads_category_if_confident, "Nieznana klasyfikacja"),127 "class_probabilities": {str(j + 1): float(mean_probabilities[j]) for j in range(len(mean_probabilities))},128 "grad_cam_image_base64": None, "error": None,129 "details": f"Uncertainty metric ({uncertainty_metric:.4f}) jest w granicach progu ({uncertainty_threshold_std}).",130 "predicted_class_index": predicted_class_index131 }132 results.append(result)133 134 return results135 136# --- Funkcja do tworzenia obrazu z nałożoną mapą Grad-CAM ---137def create_grad_cam_overlay_image(original_pil_image: Image.Image, grayscale_cam: np.ndarray, birads_category: int, transparency: float = 0.5) -> Image.Image:138 try:139 img_np = np.array(original_pil_image.convert('RGB')).astype(np.float32) / 255.0140 cam_resized = cv2.resize(grayscale_cam, (img_np.shape[1], img_np.shape[0]))141 cam_normalized = (cam_resized - np.min(cam_resized)) / (np.max(cam_resized) - np.min(cam_resized) + 1e-8)142 threshold = 0.7143 cam_normalized[cam_normalized < threshold] = 0144 kernel = np.ones((5, 5), np.uint8)145 cam_cleaned = cv2.morphologyEx(cam_normalized, cv2.MORPH_OPEN, kernel)146 birads_colors_rgb = {147 1: (0.1, 0.7, 0.1), 2: (0.53, 0.81, 0.92), 3: (1.0, 0.9, 0.0),148 4: (1.0, 0.5, 0.0), 5: (0.9, 0.1, 0.1)149 }150 chosen_color = np.array(birads_colors_rgb.get(birads_category, (0.5, 0.5, 0.5)))151 color_overlay_np = np.zeros_like(img_np)152 for c in range(3): color_overlay_np[:, :, c] = chosen_color[c]153 alpha = cam_cleaned * transparency154 alpha_expanded = alpha[..., np.newaxis]155 highlighted_image_np = img_np * (1 - alpha_expanded) + color_overlay_np * alpha_expanded156 highlighted_image_np = np.clip(highlighted_image_np, 0, 1)157 final_image_np = (highlighted_image_np * 255).astype(np.uint8)158 return Image.fromarray(final_image_np)159 except Exception as e:160 logger.error(f"Błąd podczas tworzenia obrazu Grad-CAM overlay: {e}", exc_info=True) 161 return None162 163# --- ZAKTUALIZOWANA Funkcja do heurystycznych testów OOD ---164def run_heuristic_ood_checks(pil_image: Image.Image, request_id: str, colorfulness_threshold: float, uniformity_threshold: float, aspect_ratio_min: float, aspect_ratio_max: float) -> Optional[str]:165 """166 Wykonuje heurystyki OOD. Zwraca konkretny komunikat błędu w razie problemu, w przeciwnym razie None.167 """168 logger.info(f"[RequestID: {request_id}] Uruchamianie heurystycznych testów OOD...")169 width, height = pil_image.size170 171 # Sprawdzimy najpierw kolorowość, bo to najczęstszy problem172 img_rgb_for_color_check = pil_image.convert('RGB')173 img_np_rgb = np.array(img_rgb_for_color_check)174 mean_std_across_channels = np.mean(np.std(img_np_rgb, axis=2))175 logger.info(f"[RequestID: {request_id}] Heurystyka: Kolorowość = {mean_std_across_channels:.2f} (próg: {colorfulness_threshold})")176 177 if mean_std_across_channels > colorfulness_threshold:178 # Ten komunikat jest teraz bardziej specyficzny179 msg = f"Wykryto kolorowy obraz (wskaźnik: {mean_std_across_channels:.2f}). System oczekuje obrazu w skali szarości, typowego dla badań medycznych."180 logger.warning(f"[RequestID: {request_id}] Heurystyka OOD ODRZUCONA: {msg}")181 # Zwracamy specjalny typ błędu, który potem rozpoznamy182 return f"INVALID_IMAGE_TYPE: {msg}"183 184 aspect_ratio = width / height185 if not (aspect_ratio_min < aspect_ratio < aspect_ratio_max):186 msg = f"Nietypowe proporcje obrazu: {aspect_ratio:.2f}."187 return f"HEURISTIC_FAILED: {msg}"188 189 gray_image = pil_image.convert('L')190 std_dev_intensity = np.std(np.array(gray_image))191 if std_dev_intensity < uniformity_threshold:192 msg = f"Obraz wydaje się zbyt jednolity (np. cały czarny): {std_dev_intensity:.2f}."193 return f"HEURISTIC_FAILED: {msg}"194 195 logger.info(f"[RequestID: {request_id}] Heurystyczne testy OOD zakończone pomyślnie.")196 return None197 198# --- Aplikacja FastAPI ---199class PredictionResult(BaseModel):200 birads: Optional[int] = None201 confidence: Optional[float] = None202 interpretation: str203 class_probabilities: Dict[str, float]204 grad_cam_image_base64: Optional[str] = None205 error: Optional[str] = None206 details: Optional[str] = None207 208app = FastAPI(title="BI-RADS Mammography Classification API")209 210@app.on_event("startup")211async def startup_event():212 logger.info("Rozpoczynanie eventu startup aplikacji FastAPI.")213 initialize_model()214 215# --- ZAKTUALIZOWANY Endpoint /predict/ ---216@app.post("/predict/", response_model=List[PredictionResult])217async def predict_images(218 files: List[UploadFile] = File(...),219 colorfulness_threshold: float = 2.0,220 uniformity_threshold: float = 10.0,221 aspect_ratio_min: float = 0.4,222 aspect_ratio_max: float = 2.5,223 mc_dropout_samples: int = 25,224 uncertainty_threshold_std: float = 0.11225):226 request_id = os.urandom(8).hex()227 logger.info(f"[RequestID: {request_id}] Otrzymano żądanie /predict/ dla {len(files)} plików.")228 229 if model_instance is None or transform_pipeline is None:230 raise HTTPException(status_code=503, detail="Model nie jest zainicjalizowany.")231 232 all_results = []233 valid_images_pil = []234 valid_tensors = []235 original_indices = []236 237 for idx, file in enumerate(files):238 try:239 contents = await file.read()240 image_pil_original = Image.open(io.BytesIO(contents))241 242 ood_error_details = run_heuristic_ood_checks(243 image_pil_original.copy(), request_id,244 colorfulness_threshold, uniformity_threshold, aspect_ratio_min, aspect_ratio_max245 )246 247 if ood_error_details:248 # Rozpoznajemy nasz specjalny typ błędu249 if ood_error_details.startswith("INVALID_IMAGE_TYPE"):250 error_type = "Invalid Image Type"251 interpretation = "Przesłany plik nie wygląda na obraz mammograficzny. Proszę wgrać odpowiednie zdjęcie USG."252 details = ood_error_details.replace("INVALID_IMAGE_TYPE: ", "")253 else: # Pozostałe błędy heurystyczne254 error_type = "Heuristic OOD check failed"255 interpretation = "Obraz odrzucony przez wstępne testy. Może mieć nietypowe wymiary lub być zbyt jednolity."256 details = ood_error_details.replace("HEURISTIC_FAILED: ", "")257 258 result = PredictionResult(259 interpretation=interpretation,260 class_probabilities={}, error=error_type,261 details=details262 )263 all_results.append((idx, result))264 continue265 266 image_rgb = image_pil_original.convert("RGB")267 input_tensor = transform_pipeline(image_rgb).unsqueeze(0).to(DEVICE)268 269 valid_images_pil.append(image_rgb)270 valid_tensors.append(input_tensor)271 original_indices.append(idx)272 273 except Exception as e:274 logger.error(f"[RequestID: {request_id}] Błąd podczas odczytu pliku {file.filename}: {e}", exc_info=True)275 result = PredictionResult(276 interpretation="Błąd podczas przetwarzania pliku.", class_probabilities={},277 error="File processing error.", details=str(e)278 )279 all_results.append((idx, result))280 281 if valid_tensors:282 batch_tensor = torch.cat(valid_tensors, dim=0)283 logger.info(f"[RequestID: {request_id}] Przetwarzanie wsadu {batch_tensor.shape[0]} poprawnych obrazów.")284 285 mc_results = predict_with_mc_dropout(model_instance, batch_tensor, mc_dropout_samples, uncertainty_threshold_std)286 287 model_instance.eval()288 target_layers = [model_instance.layer4[-1]]289 cam_algorithm = GradCAMPlusPlus(model=model_instance, target_layers=target_layers)290 291 for i, result_dict in enumerate(mc_results):292 if not result_dict.get("error"):293 birads_cat = result_dict["birads"]294 pred_idx = result_dict["predicted_class_index"]295 296 input_tensor_for_cam = batch_tensor[i].unsqueeze(0).clone().detach().requires_grad_(True)297 targets_for_cam = [ClassifierOutputTarget(pred_idx)]298 299 grayscale_cam = cam_algorithm(input_tensor=input_tensor_for_cam, targets=targets_for_cam)300 301 if grayscale_cam is not None:302 overlay_image_pil = create_grad_cam_overlay_image(303 original_pil_image=valid_images_pil[i],304 grayscale_cam=grayscale_cam[0, :],305 birads_category=birads_cat306 )307 if overlay_image_pil:308 buffered = io.BytesIO()309 overlay_image_pil.save(buffered, format="PNG")310 result_dict["grad_cam_image_base64"] = base64.b64encode(buffered.getvalue()).decode('utf-8')311 312 result_dict.pop("predicted_class_index", None)313 all_results.append((original_indices[i], PredictionResult(**result_dict)))314 315 all_results.sort(key=lambda x: x[0])316 final_results = [res for _, res in all_results]317 318 return final_results319 320@app.get("/")321async def root():322 logger.info("Otrzymano żądanie GET na /")323 return {"message": "Witaj w BI-RADS Classification API! Użyj endpointu /predict/ do wysyłania obrazów."}