Spaces:
Sleeping
Sleeping
| import io | |
| import numpy as np | |
| import tensorflow as tf | |
| from fastapi import FastAPI, File, UploadFile, HTTPException | |
| from PIL import Image | |
| from typing import Dict | |
| app = FastAPI(title="DenseNet121 Image Prediction API") | |
| # Labels dictionary as per requirements | |
| LABELS = { | |
| 0: {"short": "DR", "full": "Diabetic Retinopathy"}, | |
| 1: {"short": "MH", "full": "Media Haze"}, | |
| 2: {"short": "NL", "full": "Normal"} | |
| } | |
| # Global variable to store the model | |
| model = None | |
| async def load_model(): | |
| global model | |
| try: | |
| # Loading the fully trained model | |
| model = tf.keras.models.load_model("DenseNet121.h5") | |
| print("Model loaded successfully.") | |
| except Exception as e: | |
| print(f"Error loading model: {e}") | |
| def preprocess_image(image: Image.Image) -> np.ndarray: | |
| """ | |
| Preprocess the image according to requirements: | |
| - RGB format | |
| - Resize to 224x224 | |
| - Normalize by dividing by 255 | |
| """ | |
| if image.mode != "RGB": | |
| image = image.convert("RGB") | |
| image = image.resize((224, 224)) | |
| img_array = np.array(image) | |
| # Normalize by 255 | |
| img_array = img_array.astype(np.float32) / 255.0 | |
| # Add batch dimension | |
| img_array = np.expand_dims(img_array, axis=0) | |
| return img_array | |
| async def root(): | |
| return {"message": "DenseNet121 Prediction API is running", "model_status": "Loaded" if model else "Not Loaded"} | |
| async def predict(file: UploadFile = File(...)): | |
| """ | |
| Endpoint for image prediction. | |
| Accepts an image file and returns softmax probabilities and class labels. | |
| """ | |
| if model is None: | |
| raise HTTPException(status_code=500, detail="Model not loaded.") | |
| try: | |
| # Read image | |
| contents = await file.read() | |
| image = Image.open(io.BytesIO(contents)) | |
| # Preprocess | |
| processed_image = preprocess_image(image) | |
| # Predict | |
| predictions = model.predict(processed_image)[0] # Softmax probabilities | |
| # Prepare response | |
| results = {} | |
| summary_parts = [] | |
| for i, prob in enumerate(predictions): | |
| label_info = LABELS[i] | |
| percentage = prob * 100 | |
| results[label_info["full"]] = f"{percentage:.2f}%" | |
| summary_parts.append(f"{int(round(percentage))}% {label_info['short']}") | |
| summary_text = ", ".join(summary_parts) | |
| return { | |
| "prediction_summary": summary_text, | |
| "detailed_probabilities": results, | |
| "top_prediction": LABELS[np.argmax(predictions)]["full"] | |
| } | |
| except Exception as e: | |
| raise HTTPException(status_code=500, detail=f"Prediction error: {str(e)}") | |