File size: 5,010 Bytes
247540c
c31e6be
 
 
 
 
247540c
c3ba4a1
 
 
247540c
 
a847ebb
8b88424
a847ebb
 
 
 
5c6ea5f
2d916b2
496f6b6
9a782d2
2d916b2
3c483d0
 
7a4138e
 
 
 
 
 
 
 
 
c31e6be
8b88424
 
 
c3ba4a1
 
 
 
8b88424
c3ba4a1
 
8b88424
c3ba4a1
 
 
 
8b88424
 
 
3c483d0
 
8b88424
c3ba4a1
 
3c483d0
 
c3ba4a1
3c483d0
 
c31e6be
8b88424
c3ba4a1
 
 
 
 
 
 
 
 
 
 
3c483d0
 
 
c3ba4a1
 
 
 
8b88424
c3ba4a1
8b88424
c3ba4a1
8b88424
 
c3ba4a1
8b88424
 
 
 
 
 
 
 
a847ebb
8b88424
 
 
 
 
 
 
c3ba4a1
8b88424
 
a847ebb
 
c3ba4a1
 
 
 
 
 
 
 
 
 
c31e6be
8b88424
 
c31e6be
7a4138e
 
 
 
 
 
 
 
 
c3ba4a1
 
 
 
 
 
3c483d0
 
c31e6be
c3ba4a1
3c483d0
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
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()