CowCatcherAI / app.py
CowcatcherAI's picture
Update app.py
d595cfe verified
Raw History Blame Contribute Delete
3.57 kB
import gradio as gr
from ultralytics import YOLO
from PIL import Image
import os
# 1. Load Models
models = {
"CowCatcherV15": YOLO('cowcatcherV15.pt'),
"CowCatcherV16 (Experimental)": YOLO('cowcatcherV16.pt')
}
# 2. Styling & Colors
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;
font-family: 'Bebas Neue', sans-serif !important;
font-size: 1.2rem !important;
color: white !important;
}
.gradio-container label span {
font-family: 'Roboto', sans-serif !important;
font-weight: bold;
}
"""
description_text = """
**CowCatcher AI** is an open-source computer vision model designed to monitor your herd 24/7.
By analyzing footage from your barn cameras, it automatically detects "mounting" behavior โ€” the primary sign of estrus (heat).
> โš ๏ธ **Best Results: High-Angle View** โ€” Trained for security camera perspectives (4โ€“5 meters high). Eye-level photos may not work well.
"""
# 3. Prediction function
def predict(img, model_name, conf_threshold):
if img is None:
return None
results = models[model_name](img, conf=conf_threshold)
res_plotted = results[0].plot()
return Image.fromarray(res_plotted[:, :, ::-1])
# 4. Load example images
example_folder = "examples"
example_images = sorted([
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 []
# 5. Build UI
with gr.Blocks(title="CowCatcher AI") as demo:
with gr.Column(elem_classes="roboto-font"):
gr.Markdown("# ๐Ÿฎ CowCatcher AI", elem_classes="bebas-font")
gr.Markdown(description_text, elem_classes="roboto-font")
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="CowCatcherV15",
label="Select Model"
)
conf_slider = gr.Slider(
minimum=0.0,
maximum=1.0,
value=0.7,
step=0.05,
label="Confidence Threshold",
info="Higher = fewer but more certain detections."
)
predict_btn = gr.Button("๐Ÿ” Run Detection", variant="primary", elem_classes="primary-btn")
with gr.Column():
output_img = gr.Image(type="pil", label="Detection Result")
if example_images:
gr.Markdown("### ๐Ÿ“ธ Try an Example", elem_classes="bebas-font")
gr.Examples(
examples=example_images,
inputs=input_img,
label=None
)
predict_btn.click(
fn=predict,
inputs=[input_img, model_drop, conf_slider],
outputs=output_img
)
if __name__ == "__main__":
demo.launch(
css=custom_css,
theme=gr.themes.Soft(primary_hue="green"),
)