Thesis / app.py
006f86's picture
added text
e5f266f
Raw
History Blame Contribute Delete
3.96 kB
import gradio as gr
import torch
import torch.nn as nn
from torchvision import transforms
from torchvision.models import efficientnet_v2_s, EfficientNet_V2_S_Weights
from PIL import Image
import numpy as np
import cv2
from ultralytics import YOLO
import json
import os
# ---------------------------
# 1. Load EfficientNet model
# ---------------------------
MODEL_PATH = "best_model.pth"
with open("class_names.json", "r") as f:
CLASS_NAMES = json.load(f)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
weights = EfficientNet_V2_S_Weights.IMAGENET1K_V1
model = efficientnet_v2_s(weights=weights)
model.classifier[1] = nn.Linear(model.classifier[1].in_features, len(CLASS_NAMES))
model.load_state_dict(torch.load(MODEL_PATH, map_location=device))
model.eval().to(device)
mean = getattr(weights, "meta", {}).get("mean", [0.485, 0.456, 0.406])
std = getattr(weights, "meta", {}).get("std", [0.229, 0.224, 0.225])
transform = transforms.Compose([
transforms.Resize((384, 384)),
transforms.ToTensor(),
transforms.Normalize(mean=mean, std=std),
])
# ---------------------------
# 2. Load YOLOv8 model
# ---------------------------
yolo_model = YOLO("yolov8x.pt")
# ---------------------------
# 3. Detection + Identification Pipeline
# ---------------------------
def detect_and_identify(image: Image.Image):
# Convert PIL to OpenCV (numpy)
image_cv = cv2.cvtColor(np.array(image), cv2.COLOR_RGB2BGR)
# Run YOLOv8 detection
results = yolo_model(image_cv)
boxes = results[0].boxes
names = results[0].names
detections = boxes.data.cpu().numpy()
# Filter for cats only
valid_detections = []
for det in detections:
x1, y1, x2, y2, conf, cls = det
class_name = names[int(cls)]
if class_name == "cat":
valid_detections.append(det)
if not valid_detections:
return (
None,
{"Error": "No cat detected in the image."},
{"Info": "Please upload an image containing one or more visible cats."}
)
cropped_images = []
predictions = {}
for i, det in enumerate(valid_detections):
x1, y1, x2, y2, conf, cls = det
x1, y1, x2, y2 = map(int, [x1, y1, x2, y2])
# Crop the detected cat
cropped_cat = image_cv[y1:y2, x1:x2]
cropped_pil = Image.fromarray(cv2.cvtColor(cropped_cat, cv2.COLOR_BGR2RGB))
# Preprocess for classifier
input_tensor = transform(cropped_pil).unsqueeze(0).to(device)
# Identify using EfficientNet
with torch.no_grad():
outputs = model(input_tensor)
probs = torch.nn.functional.softmax(outputs, dim=1)[0]
# Get top prediction
top_idx = torch.argmax(probs).item()
top_label = CLASS_NAMES[top_idx]
confidence = float(probs[top_idx])
# Get Top 5 predictions only
top5_indices = torch.argsort(probs, descending=True)[:5]
top5_results = {CLASS_NAMES[j]: float(probs[j]) for j in top5_indices}
# Append image with caption (for Gallery)
cropped_images.append((cropped_pil, f"{top_label} ({confidence*100:.1f}%)"))
# Add JSON data for this cat
predictions[f"Cat {i+1} - {top_label}"] = {
"Top Prediction": top_label,
"Confidence": confidence,
"Top 5 Scores": top5_results
}
return cropped_images, predictions
# ---------------------------
# 4. Gradio Interface
# ---------------------------
demo = gr.Interface(
fn=detect_and_identify,
inputs=gr.Image(type="pil", label="Upload Image"),
outputs=[
gr.Gallery(label="Detected Cats (Cropped)", columns=2, rows=2),
gr.JSON(label="Predictions per Cat")
],
title="StrAI - Cat Identifier - 12/2/2025 Build",
description="Upload an image with one or more cats.",
api_name="predict"
)
if __name__ == "__main__":
demo.launch(server_name="0.0.0.0")