File size: 2,637 Bytes
a31ac05
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
194a6df
a31ac05
842a0b3
a31ac05
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
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()