SabaTariq510's picture
Update app.py
81dca0b verified
Raw
History Blame Contribute Delete
5.35 kB
import os
# Reduce TensorFlow log verbosity
os.environ["TF_CPP_MIN_LOG_LEVEL"] = "2"
import numpy as np
import torch
import gradio as gr
import spaces
from huggingface_hub import hf_hub_download
from tensorflow.keras.models import load_model
from tensorflow.keras.preprocessing import image as keras_image
from ultralytics import YOLO
# =========================================================
# Your model repo (where you uploaded the trained models)
# =========================================================
MODEL_REPO = "SabaTariq510/waste-classification-models"
# =========================================================
# Download + load models (runs once when the Space starts)
# =========================================================
print("Downloading MobileNetV2...")
mobilenet_path = hf_hub_download(repo_id=MODEL_REPO, filename="best_mobilenet.keras")
print("Downloading YOLOv8...")
yolo_path = hf_hub_download(repo_id=MODEL_REPO, filename="bestyolomodel.pt")
print("Loading MobileNetV2...")
mobilenet_model = load_model(mobilenet_path)
print("Loading YOLOv8...")
yolo_model = YOLO(yolo_path, task="detect")
print("Both models loaded successfully!")
# =========================================================
# MobileNet Classes
# =========================================================
mobilenet_classes = [
"cardboard",
"glass",
"metal",
"paper",
"plastic",
"trash"
]
# =========================================================
# Recyclable Information
# =========================================================
recyclable = {
"cardboard": "Recyclable",
"glass": "Recyclable",
"metal": "Recyclable",
"paper": "Recyclable",
"plastic": "Recyclable",
"trash": "Non-Recyclable"
}
# =========================================================
# Prediction Function
# ZeroGPU: this function must be decorated with @spaces.GPU
# for it to actually get a GPU allocated to it, otherwise
# no GPU will be assigned at runtime.
# =========================================================
@spaces.GPU
def predict(img, model_choice):
if img is None:
return "⚠️ Please upload an image first."
# Save a temporary copy (YOLO needs a file path as input)
temp_path = "/tmp/uploaded_image.jpg"
img.save(temp_path)
# =====================================================
# MobileNetV2
# =====================================================
if model_choice == "MobileNetV2":
resized = img.convert("RGB").resize((224, 224))
arr = keras_image.img_to_array(resized)
arr = arr / 255.0
arr = np.expand_dims(arr, axis=0)
prediction = mobilenet_model.predict(arr, verbose=0)
index = int(np.argmax(prediction))
confidence = float(np.max(prediction)) * 100
waste = mobilenet_classes[index]
model_used = "MobileNetV2"
# =====================================================
# YOLOv8
# =====================================================
else:
# Only use "cuda" if a GPU is actually available at runtime
# (e.g. ZeroGPU Space). Otherwise fall back to CPU so the
# prediction doesn't silently fail with CUDA_ERROR_NO_DEVICE.
device = "cuda" if torch.cuda.is_available() else "cpu"
try:
results = yolo_model.predict(
source=temp_path,
conf=0.25,
device=device
)
except Exception as e:
return f"❌ YOLOv8 prediction failed: {e}"
result = results[0]
if len(result.boxes) > 0:
# Confidence values for all detected boxes
confidences = result.boxes.conf
# Pick the detection with the highest confidence
best_index = int(confidences.argmax())
box = result.boxes[best_index]
class_index = int(box.cls)
confidence = float(box.conf) * 100
waste = result.names[class_index]
else:
waste = "No Waste Detected"
confidence = 0
model_used = f"YOLOv8 ({device})"
recycle = recyclable.get(waste, "Unknown")
result_text = (
f"**Model Used:** {model_used}\n\n"
f"**Prediction:** {waste}\n\n"
f"**Confidence:** {confidence:.2f}%\n\n"
f"**Recyclable:** {recycle}"
)
return result_text
# =========================================================
# Gradio Interface
# =========================================================
with gr.Blocks(title="Smart Waste Classification") as demo:
gr.Markdown("# ♻️ Smart Waste Classification")
gr.Markdown("Upload an image and choose a model — get an instant prediction.")
with gr.Row():
with gr.Column():
img_input = gr.Image(type="pil", label="Upload Waste Image")
model_choice = gr.Radio(
choices=["MobileNetV2", "YOLOv8"],
value="MobileNetV2",
label="Select Model"
)
predict_btn = gr.Button("Predict", variant="primary")
with gr.Column():
output = gr.Markdown(label="Result")
predict_btn.click(
fn=predict,
inputs=[img_input, model_choice],
outputs=output
)
if __name__ == "__main__":
demo.launch()