CoolFace
Apppublic

Enterwar99/mammography-api-server

sourceHugging Faceapache-2.0updated 11mo agoView on Hugging Face
0likes
api_app.py323 linesDownload Raw Back to root
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."}