Object_Counter / app.py
Harbidel's picture
Upload app.py
8e60403 verified
Raw History Blame Contribute Delete
10.3 kB
"""
Object Counter β€” Gradio App (Python script version)
Detects and counts any object(s) a user names in an uploaded image, using
YOLO-World (open-vocabulary object detection). No fixed class list β€” type
what you want counted at runtime.
Usage:
pip install -r requirements.txt
python app.py
Then open the local URL Gradio prints (usually http://127.0.0.1:7860).
"""
from collections import Counter
import gradio as gr
import supervision as sv
import torch
from ultralytics import YOLOWorld
# ---------------------------------------------------------------------------
# Workaround for a known, still-open Gradio bug (gradio-app/gradio#11722):
# Gradio's automatic API-schema generation crashes with
# "TypeError: argument of type 'bool' is not iterable" when a component's
# JSON schema contains a boolean where a dict is expected. This has hit many
# projects across several Gradio versions and isn't reliably fixed by
# pinning one version, so we patch the offending function defensively.
# This only affects auto-generated API docs, not the app's own behavior.
# ---------------------------------------------------------------------------
def _patch_gradio_bool_schema_bug():
try:
import gradio_client.utils as client_utils
_original_get_type = client_utils.get_type
def _patched_get_type(schema):
if isinstance(schema, bool):
return "bool"
return _original_get_type(schema)
client_utils.get_type = _patched_get_type
if hasattr(client_utils, "_json_schema_to_python_type"):
_original_json_type = client_utils._json_schema_to_python_type
def _patched_json_type(schema, defs=None):
if isinstance(schema, bool):
return "bool"
return _original_json_type(schema, defs)
client_utils._json_schema_to_python_type = _patched_json_type
except Exception:
# If gradio_client's internals change shape in a future version,
# fail quietly rather than blocking the app from starting.
pass
_patch_gradio_bool_schema_bug()
# ---------------------------------------------------------------------------
# Model setup
# ---------------------------------------------------------------------------
# "yolov8s-worldv2.pt" is a good balance of speed/accuracy. Swap for
# "yolov8m-worldv2.pt" or "yolov8l-worldv2.pt" if you want more accuracy
# and have the GPU/CPU headroom for it. Weights download automatically
# on first run.
MODEL_WEIGHTS = "yolov8s-worldv2.pt"
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
model = YOLOWorld(MODEL_WEIGHTS)
model.to(DEVICE)
box_annotator = sv.BoxAnnotator(thickness=2)
label_annotator = sv.LabelAnnotator(text_scale=0.5, text_thickness=1)
EXAMPLE_PROMPTS = [
"person, car",
"car, truck, bus, bicycle",
"steel rod, stick",
"bottle, cup, chair",
]
# ---------------------------------------------------------------------------
# Core detection + counting logic
# ---------------------------------------------------------------------------
def count_objects(image, class_text, confidence):
"""
Run open-vocabulary detection on `image` for the classes in `class_text`
(comma-separated). Returns:
- the annotated image
- a [ [class, count], ... ] table for the results grid
- an HTML string with a big total-count readout
- a status message
"""
if image is None:
return None, [], _total_html(None), "⚠️ Please upload an image first."
if not class_text or not class_text.strip():
return (
None,
[],
_total_html(None),
"⚠️ Type at least one object to count, e.g. 'car, person'.",
)
# Parse comma-separated class names, strip whitespace, drop empties
classes = [c.strip() for c in class_text.split(",") if c.strip()]
# Workaround for a known ultralytics/CLIP device-mismatch bug: set_classes()
# loads a fresh CLIP text encoder that can end up on a different device than
# the rest of the model when running on GPU. Moving to CPU for this call,
# then back to the target device, avoids the crash.
model.to("cpu")
model.set_classes(classes)
model.to(DEVICE)
# Gradio passes images in as RGB numpy arrays already
results = model.predict(image, conf=confidence, verbose=False)
result = results[0]
if len(result.boxes) == 0:
table = [[cls, 0] for cls in classes]
status = (
"No objects detected. Try lowering the confidence threshold "
"or rewording the class names."
)
return image, table, _total_html(0), status
detections = sv.Detections.from_ultralytics(result)
class_ids = result.boxes.cls.cpu().numpy().astype(int)
class_names = [result.names[i] for i in class_ids]
counts = Counter(class_names)
labels = [
f"{result.names[cid]} {conf:.2f}"
for cid, conf in zip(detections.class_id, detections.confidence)
]
annotated = box_annotator.annotate(scene=image.copy(), detections=detections)
annotated = label_annotator.annotate(scene=annotated, detections=detections, labels=labels)
# Table rows, in the order the user typed the classes
table = [[cls, counts.get(cls, 0)] for cls in classes]
total = sum(counts.values())
status = f"βœ… Done β€” found {total} object{'s' if total != 1 else ''}."
return annotated, table, _total_html(total), status
def _total_html(total):
"""Small styled readout for the total object count."""
display_value = "β€”" if total is None else str(total)
return (
"<div style='text-align:center; padding: 6px 0;'>"
"<div style='font-size:0.85rem; letter-spacing:0.05em; text-transform:uppercase; "
"color:var(--body-text-color-subdued); margin-bottom:2px;'>Total objects</div>"
f"<div style='font-size:2.4rem; font-weight:700; line-height:1;'>{display_value}</div>"
"</div>"
)
def clear_all():
return None, "", 0.15, None, [], _total_html(None), ""
# ---------------------------------------------------------------------------
# Gradio UI
# ---------------------------------------------------------------------------
CUSTOM_CSS = """
#header { text-align: center; margin-bottom: 4px; }
#header h1 { margin-bottom: 2px; }
#header p { color: var(--body-text-color-subdued); margin-top: 0; }
#run_btn { min-height: 46px; font-size: 1.05rem; }
#results_table table { font-size: 0.95rem; }
.gr-accordion { margin-top: 8px; }
"""
THEME = gr.themes.Soft(
primary_hue="blue",
secondary_hue="slate",
neutral_hue="slate",
font=[gr.themes.GoogleFont("Inter"), "ui-sans-serif", "system-ui", "sans-serif"],
)
def build_app():
with gr.Blocks(title="Object Counter", theme=THEME, css=CUSTOM_CSS) as demo:
with gr.Column(elem_id="header"):
gr.Markdown(
"# πŸ”Ž Object Counter\n"
"Upload a photo, tell it what to look for, and get a count β€” "
"for **any object**, not just a fixed list."
)
with gr.Row(equal_height=False):
# -------------------- Left: inputs --------------------
with gr.Column(scale=1):
image_input = gr.Image(
label="Image",
type="numpy",
height=340,
)
class_input = gr.Textbox(
label="What should I count?",
placeholder="e.g. car, person, steel rod",
info="Separate multiple objects with commas.",
)
gr.Examples(
examples=EXAMPLE_PROMPTS,
inputs=class_input,
label="Quick examples",
)
confidence_input = gr.Slider(
minimum=0.01,
maximum=0.9,
value=0.15,
step=0.01,
label="Confidence threshold",
info="Lower catches more objects but risks false positives; higher is stricter.",
)
with gr.Row():
clear_btn = gr.Button("Clear", scale=1)
submit_btn = gr.Button("Count objects", variant="primary", scale=2, elem_id="run_btn")
with gr.Accordion("πŸ’‘ Tips for better results", open=False):
gr.Markdown(
"- Be descriptive: **\"steel rod\"** works better than **\"metal\"**.\n"
"- Missing objects? Lower the confidence threshold.\n"
"- Too many false positives? Raise the threshold or narrow the wording.\n"
"- Objects that are tightly packed or overlapping (e.g. a bundle of rods) "
"may be undercounted β€” this is a limitation of box-based detection."
)
# -------------------- Right: outputs --------------------
with gr.Column(scale=1):
image_output = gr.Image(label="Detected objects", height=340)
total_output = gr.HTML(_total_html(None))
results_table = gr.Dataframe(
headers=["Object", "Count"],
datatype=["str", "number"],
row_count=(0, "dynamic"),
col_count=(2, "fixed"),
label="Breakdown",
elem_id="results_table",
)
status_output = gr.Markdown("")
submit_btn.click(
fn=count_objects,
inputs=[image_input, class_input, confidence_input],
outputs=[image_output, results_table, total_output, status_output],
)
clear_btn.click(
fn=clear_all,
inputs=None,
outputs=[
image_input,
class_input,
confidence_input,
image_output,
results_table,
total_output,
status_output,
],
)
return demo
if __name__ == "__main__":
app = build_app()
app.launch(server_name="0.0.0.0", server_port=7860, ssr_mode=False)