CowcatcherAI's picture
Update app.py
a847ebb verified
Raw History Blame Contribute Delete
5.01 kB
import os
import platform
import pathlib
import gradio as gr
from PIL import Image
from ultralytics import YOLO
# Fix for Windows paths on Linux/Mac environments
plt_sys = platform.system()
if plt_sys != 'Windows':
pathlib.WindowsPath = pathlib.PosixPath
# Load Models (V7 t/m V10 toegevoegd)
MODELS = {
"CalvingCatcherV10": YOLO("calvingcatcherV10.pt"),
"CalvingCatcherV9": YOLO("calvingcatcherV9.pt"),
"CalvingCatcherV8": YOLO("calvingcatcherV8.pt"),
"CalvingCatcherV7": YOLO("calvingcatcherV7.pt"),
"CalvingCatcherV6": YOLO("calvingcatcherV6.pt"),
"CalvingCatcherV5": YOLO("calvingcatcherV5.pt"),
"CalvingCatcherV4": YOLO("calvingcatcherV4.pt"),
"CalvingCatcherV3": YOLO("calvingcatcherV3.pt"),
"CalvingCatcherV2": YOLO("calvingcatcherV2.pt")
}
# --- Setup Example Images ---
example_folder = "examples"
example_images = [
os.path.join(example_folder, f)
for f in os.listdir(example_folder)
if f.lower().endswith(('.png', '.jpg', '.jpeg'))
] if os.path.exists(example_folder) else []
# ----------------------------
CUSTOM_CSS = """
@import url('https://fonts.googleapis.com/css2?family=Roboto:wght@400;700&display=swap');
@import url('https://fonts.googleapis.com/css2?family=Bebas+Neue&display=swap');
.bebas-font {
font-family: 'Bebas Neue', sans-serif !important;
text-transform: uppercase;
letter-spacing: 1px;
}
.roboto-font {
font-family: 'Roboto', sans-serif !important;
}
.primary-btn {
background-color: #386938 !important;
border: none !important;
color: white !important;
font-family: 'Bebas Neue', sans-serif !important;
font-size: 1.2rem !important;
}
"""
DESCRIPTION_TEXT = """
**CalvingCatcher AI** is an informative open-source computer vision tool designed to monitor calving processes 24/7.
It analyzes footage to detect critical signs and stages of birth in cattle.
"""
def predict(img, model_name, conf_threshold, selected_class_names):
if img is None:
return None
selected_model = MODELS[model_name]
model_classes = selected_model.names
selected_indices = [
idx for idx, name in model_classes.items()
if name in selected_class_names
]
if not selected_indices:
return img
results = selected_model(img, conf=conf_threshold, classes=selected_indices)
res_plotted = results[0].plot()
return Image.fromarray(res_plotted[:, :, ::-1])
def update_checkboxes(model_name):
model = MODELS[model_name]
names = list(model.names.values())
return gr.update(choices=names, value=names)
# Build Gradio Interface
with gr.Blocks(css=CUSTOM_CSS, theme=gr.themes.Soft(primary_hue="green"), title="CalvingCatcher AI") as demo:
with gr.Column(elem_classes="roboto-font"):
gr.Markdown("# 🐮 CalvingCatcher AI", elem_classes="bebas-font")
gr.Markdown(DESCRIPTION_TEXT)
with gr.Row():
with gr.Column():
input_img = gr.Image(type="pil", label="Upload Barn Image")
with gr.Row():
model_drop = gr.Dropdown(
choices=list(MODELS.keys()),
value="CalvingCatcherV10", # Aangepast naar V10
label="Select Model"
)
conf_slider = gr.Slider(
minimum=0.1,
maximum=1.0,
value=0.3,
label="Confidence Threshold",
info="Higher values = more certainty, fewer detections."
)
# initial_model aangepast naar V10 zodat de checkboxes direct kloppen
initial_model = "CalvingCatcherV10"
initial_choices = list(MODELS[initial_model].names.values())
class_selector = gr.CheckboxGroup(
choices=initial_choices,
value=initial_choices,
label="Filter Detection Classes",
info="Uncheck classes you want to ignore."
)
predict_btn = gr.Button("Run Detection", variant="primary", elem_classes="primary-btn")
with gr.Column():
output_img = gr.Image(type="pil", label="Detection Result")
# --- Display Examples if they exist ---
if example_images:
gr.Markdown("### 📸 Try an Example:", elem_classes="bebas-font")
gr.Examples(
examples=example_images,
inputs=input_img,
label=None
)
# Event Listeners
model_drop.change(
fn=update_checkboxes,
inputs=model_drop,
outputs=class_selector
)
predict_btn.click(
fn=predict,
inputs=[input_img, model_drop, conf_slider, class_selector],
outputs=output_img
)
if __name__ == "__main__":
demo.launch()