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