Mohith202's picture
Update app.py
194a6df verified
Raw
History Blame Contribute Delete
2.64 kB
import torch
from torchvision import transforms
import gradio as gr
import cv2
import numpy as np
from segmentation_models_pytorch.unetplusplus.model import UnetPlusPlus
import matplotlib.pyplot as plt
from PIL import Image
# Define the function to load the model
def get_seg_model(candidate, weight_path):
model = UnetPlusPlus(
encoder_name=candidate['backbone_name'],
encoder_depth=5,
encoder_weights=None,
classes=2, # Output 2 classes
activation='sigmoid',
)
model.load_state_dict(torch.load(weight_path, map_location='cpu'))
return model
# Load the segmentation model
SEG_MODEL = {
'backbone_name': 'densenet121',
'pretranied_weight': 'Fold0_densenet121_2d_segment (1).pth'
}
seg_model = get_seg_model(SEG_MODEL, SEG_MODEL['pretranied_weight'])
seg_model.eval() # Set the model to evaluation mode
# Preprocessing function
def preprocess_image(image):
# Convert image to numpy array and process
np_image = np.array(image)
# If image has 3 channels, transpose it to (C, H, W)
if np_image.ndim == 3:
np_image = np_image.transpose(2, 0, 1)
elif np_image.ndim == 2:
np_image = np.expand_dims(np_image, axis=0) # For grayscale (1, H, W)
np_image = np_image.astype(np.float32) / 255.0 # Scale to [0, 1]
# Normalize based on ImageNet stats (for DenseNet)
mean = np.array([0.485, 0.456, 0.406]).reshape(3, 1, 1)
std = np.array([0.229, 0.224, 0.225]).reshape(3, 1, 1)
np_image = (np_image - mean) / std
# Convert to tensor and add batch dimension (Shape: (1, C, H, W))
tensor_image = torch.tensor(np_image, dtype=torch.float32).unsqueeze(0)
# Run the model on the image
with torch.no_grad():
output = seg_model(tensor_image)
return output
# Function to generate the segmentation mask for Gradio output
def generate_output(image):
# Preprocess image and get the segmentation output
output = preprocess_image(image)
# Extract the segmentation masks (assuming binary mask, 2 classes)
mask_class_1 = output[0, 0, :, :].detach().numpy() # First class mask
mask_class_2 = output[0, 1, :, :].detach().numpy() # Second class mask
# Stack the masks together and return as two separate outputs for Gradio
return mask_class_1, mask_class_2
# Create the Gradio interface
interface = gr.Interface(
fn=generate_output, # Function to call on image input
inputs=gr.Image(type="pil"), # Input image as PIL
outputs=[gr.Image(), gr.Image()] # Two outputs for two masks
)
# Launch the interface
interface.launch()