File size: 10,337 Bytes
6d0e7a4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8e60403
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6d0e7a4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8e60403
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
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
"""
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)