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()