wng-mchl
changed the confidence score output format to percentage
ea6b566
Raw
History Blame Contribute Delete
4.89 kB
import gradio as gr
import numpy as np
import torch
import torch.nn.functional as F
from torchvision import transforms
from torchvision.models import mobilenet_v2
import tensorflow as tf
from tensorflow.keras.applications.efficientnet import preprocess_input
import matplotlib.pyplot as plt
import io
from PIL import Image
#
# --- Class Names ---
class_names = ['bidayuh', 'iban', 'kadazandusun', 'orang ulu']
# --- Load TensorFlow (EfficientNet) Model ---
efficientnet_model = tf.keras.models.load_model('models/efficientnet.keras')
# --- Load TensorFlow (Custom CNN) ---
custom_cnn_model = tf.keras.models.load_model('models/custom_cnn_model.keras')
# --- Load PyTorch (MobileNet) Model ---
torch_model = mobilenet_v2(weights=False)
torch_model.classifier[1] = torch.nn.Linear(torch_model.last_channel, len(class_names))
torch_model.load_state_dict(torch.load('models/final_best_mobilenetv2_model.pth', map_location='cpu'))
torch_model.eval()
# --- PyTorch Transform ---
torch_transform = transforms.Compose([
transforms.Resize((224, 224)),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
# --- Convert Matplotlib Figure to PIL Image ---
def fig_to_pil(fig):
buf = io.BytesIO()
fig.savefig(buf, format='png', bbox_inches='tight')
plt.close(fig)
buf.seek(0)
return Image.open(buf)
# --- TensorFlow Preprocessing ---
def preprocess_tf_image(img):
img_resized = img.resize((224, 224)).convert('RGB')
img_array = preprocess_input(np.array(img_resized))
return np.expand_dims(img_array, axis=0)
# --- Horizontal Bar Chart Generator ---
def create_prediction_chart(class_names, confidences):
fig, ax = plt.subplots(figsize=(6, 3))
bars = ax.barh(class_names, confidences, color='skyblue')
ax.set_xlim([0, 1])
ax.set_xlabel("Confidence")
ax.set_title("Class Probabilities")
# Add values on bars
for bar, confidence in zip(bars, confidences):
width = bar.get_width()
ax.text(width + 0.01, bar.get_y() + bar.get_height()/2, f"{confidence*100:.2f}%",
va='center', fontsize=9, fontweight='bold')
plt.tight_layout()
return fig_to_pil(fig)
# --- TensorFlow (EfficientNet) Prediction ---
def predict_with_efficientnet(img):
img_array = preprocess_tf_image(img)
pred = efficientnet_model.predict(img_array, verbose=0)[0]
top_idx = int(np.argmax(pred))
caption = f"{class_names[top_idx]} ({pred[top_idx]*100:.2f}%)"
chart = create_prediction_chart(class_names, pred.tolist())
return (img, caption), chart
# --- TensorFlow (Custom CNN) Prediction ---
def predict_with_custom_cnn(img):
img_array = preprocess_tf_image(img)
pred = custom_cnn_model.predict(img_array, verbose=0)[0]
top_idx = int(np.argmax(pred))
caption = f"{class_names[top_idx]} ({pred[top_idx] * 100:.2f}%)"
chart = create_prediction_chart(class_names, pred.tolist())
return (img, caption), chart
# --- PyTorch Prediction ---
def predict_with_torch(img):
input_tensor = torch_transform(img).unsqueeze(0)
with torch.no_grad():
output = torch_model(input_tensor)
probs_tensor = F.softmax(output[0], dim=0)
probs_list = probs_tensor.tolist()
top_idx = int(torch.argmax(probs_tensor))
caption = f"{class_names[top_idx]} ({probs_list[top_idx] * 100:.2f}%)"
chart = create_prediction_chart(class_names, probs_list)
return (img, caption), chart
# --- Combined Multi-Image Prediction ---
def predict_multiple_images(image_list, model_choice):
images = []
charts = []
for img in image_list:
try:
img = Image.open(img).convert('RGB')
if model_choice == "EfficientNet (TensorFlow)":
image_out, chart_out = predict_with_efficientnet(img)
elif model_choice == "Custom CNN (TensorFlow)":
image_out, chart_out = predict_with_custom_cnn(img)
elif model_choice == "MobileNet (PyTorch)":
image_out, chart_out = predict_with_torch(img)
else:
continue
images.append(image_out)
charts.append(chart_out)
except Exception as e:
print(f"Error processing image: {e}")
return images, charts
# --- Gradio Interface ---
demo = gr.Interface(
fn=predict_multiple_images,
inputs=[
gr.Files(file_types=['image']),
gr.Radio(["EfficientNet (TensorFlow)","Custom CNN (TensorFlow)", "MobileNet (PyTorch)"], label="Choose Model", value="EfficientNet (TensorFlow)")
],
outputs=[
gr.Gallery(label="Predictions with Labels"),
gr.Gallery(label="Prediction Charts")
],
title="Multi-Image Classifier (TF + PyTorch)",
description="Upload multiple images and choose a model to see predictions with confidence bar charts."
)
demo.launch(debug=True)