nitinvig's picture
Create app.py
3230905 verified
Raw
History Blame Contribute Delete
4.36 kB
import gradio as gr
import torch
from torchvision import transforms
from model import UNet
from PIL import Image
import numpy as np
# --- Configuration ---
MODES = [
"MP+Tr+BCE",
"MP+Tr+Dice",
"StrConv+Tr+BCE",
"StrConv+Ups+Dice"
]
def get_config(mode_name):
if mode_name == "MP+Tr+BCE":
return {"downsample": "maxpool", "upsample": "transpose", "filename": "unet_MP_Tr_BCE.pth"}
elif mode_name == "MP+Tr+Dice":
return {"downsample": "maxpool", "upsample": "transpose", "filename": "unet_MP_Tr_Dice.pth"}
elif mode_name == "StrConv+Tr+BCE":
return {"downsample": "strided", "upsample": "transpose", "filename": "unet_StrConv_Tr_BCE.pth"}
elif mode_name == "StrConv+Ups+Dice":
return {"downsample": "strided", "upsample": "upsample", "filename": "unet_StrConv_Ups_Dice.pth"}
else:
raise ValueError(f"Unknown mode: {mode_name}")
# --- Load Models Once (Optional optimization, but safer to load on demand if memory is tight) ---
# For immediate responsiveness, let's load them on demand inside the functions or keep them cached.
# Given it's a demo, loading 4 models might be heavy on CPU/RAM if hosted on free tier.
# But let's try to run them sequentially.
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
def run_inference_single_model(model, image_tensor):
with torch.no_grad():
output = model(image_tensor)
output = torch.sigmoid(output)
output = output.squeeze().cpu().numpy()
mask = (output > 0.5).astype(np.float32)
return mask
def predict_all(image):
if image is None:
return [None] * 4
# Preprocess
transform = transforms.Compose([
transforms.Resize((128, 128)),
transforms.ToTensor()
])
input_tensor = transform(image).unsqueeze(0).to(device)
# Create overlay
# Transform image to numpy for blending. We need it in [0,1] or [0,255]
img_resized = image.resize((128, 128))
img_np = np.array(img_resized)
results = []
for mode in MODES:
config = get_config(mode)
# Initialize model structure
model = UNet(n_channels=3, n_classes=1,
downsample_mode=config['downsample'],
upsample_mode=config['upsample']).to(device)
# Load weights
try:
model.load_state_dict(torch.load(config['filename'], map_location=device))
model.eval()
mask = run_inference_single_model(model, input_tensor)
# Create Overlay
# Mask is (128, 128). We want to make it Red where mask is 1.
# Convert mask to RGBA
overlay = np.zeros_like(img_np)
# Set Red channel to 255 where mask is 1
overlay[:,:,0] = mask * 255
# overlay[:,:,1] = 0
# overlay[:,:,2] = 0
# Blend: 0.7 * Original + 0.3 * Red Overlay (where mask is present)
# Actually simplest is just simpler blending
# Where mask == 1, add red tint
# Let's perform a weighted add using OpenCV logic manually
blended = img_np.copy()
# Indices where mask is active
idx = (mask > 0)
# Increase Red component, decrease Green/Blue to make it pop
blended[idx, 0] = np.clip(blended[idx, 0] * 0.5 + 255 * 0.5, 0, 255)
blended[idx, 1] = blended[idx, 1] * 0.5
blended[idx, 2] = blended[idx, 2] * 0.5
results.append(blended)
except Exception as e:
print(f"Error for {mode}: {e}")
results.append(img_np) # Return original image on error
return results
# --- Gradio App ---
title = "UNet Oxford-IIIT Pet Separation - Model Comparison"
description = "Compare the segmentation results of 4 different UNet configurations on the same input image."
iface = gr.Interface(
fn=predict_all,
inputs=gr.Image(type="pil", label="Input Image"),
outputs=[
gr.Image(label="MP + Tr + BCE"),
gr.Image(label="MP + Tr + Dice"),
gr.Image(label="StrConv + Tr + BCE"),
gr.Image(label="StrConv + Ups + Dice")
],
title=title,
description=description
)
if __name__ == "__main__":
iface.launch()