Add project files
This commit is contained in:
94
pipeline/yolo_world.py
Normal file
94
pipeline/yolo_world.py
Normal file
@@ -0,0 +1,94 @@
|
||||
"""
|
||||
Open-vocabulary object detection via YOLO-World (YOLOv8 family, ultralytics),
|
||||
run alongside SAM (see pipeline/flower_count.py) for the flower-count feature.
|
||||
|
||||
SAM has no notion of "flower" -- it blindly proposes every objectlike region
|
||||
it can find, so the vase and any ribbon/bow tied around it get segmented and
|
||||
counted as if they were flowers. YOLO-World, prompted with plain-language
|
||||
classes, actually knows what those things look like: its "vase"/"ribbon"/
|
||||
"bow" detections are used (in engine.count_flowers) to exclude those regions
|
||||
from SAM's flower-instance count.
|
||||
|
||||
Its own "flower" boxes are shown to the user as a second, independent count
|
||||
-- in practice it tends to draw one box per contiguous flower region rather
|
||||
than per individual bloom, so it's a coarser, corroborating signal, not a
|
||||
replacement for SAM's per-instance count.
|
||||
"""
|
||||
|
||||
import gc
|
||||
import logging
|
||||
|
||||
import cv2
|
||||
import torch
|
||||
|
||||
import config
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
_model = None
|
||||
|
||||
|
||||
def get_model():
|
||||
global _model
|
||||
if _model is None:
|
||||
logger.info("Loading YOLO-World (%s) on %s...", config.YOLO_WORLD_CHECKPOINT, DEVICE)
|
||||
from ultralytics import YOLO
|
||||
_model = YOLO(config.YOLO_WORLD_CHECKPOINT)
|
||||
_model.set_classes(config.YOLO_WORLD_CLASSES)
|
||||
return _model
|
||||
|
||||
|
||||
def unload_model():
|
||||
global _model
|
||||
freed = _model is not None
|
||||
_model = None
|
||||
if freed:
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
return freed
|
||||
|
||||
|
||||
def detect(bgr):
|
||||
"""Returns {class_name: [{"xyxy": [x1,y1,x2,y2], "confidence": float}, ...]}
|
||||
for every class in config.YOLO_WORLD_CLASSES (empty list if none found)."""
|
||||
model = get_model()
|
||||
with torch.no_grad():
|
||||
results = model.predict(bgr, conf=config.YOLO_WORLD_CONF,
|
||||
iou=config.YOLO_WORLD_IOU, verbose=False)
|
||||
r = results[0]
|
||||
by_class = {name: [] for name in config.YOLO_WORLD_CLASSES}
|
||||
for b in r.boxes:
|
||||
name = r.names[int(b.cls[0])]
|
||||
by_class.setdefault(name, []).append({
|
||||
"xyxy": [float(v) for v in b.xyxy[0].tolist()],
|
||||
"confidence": round(float(b.conf[0]), 3),
|
||||
})
|
||||
return by_class
|
||||
|
||||
|
||||
_BOX_COLORS = {
|
||||
"flower": (66, 133, 244), "vase": (219, 68, 55),
|
||||
"ribbon": (244, 180, 0), "bow": (15, 157, 88),
|
||||
}
|
||||
|
||||
|
||||
def render_boxes(bgr, by_class):
|
||||
"""Draws a labeled box per detection, colored by class, so the user can
|
||||
see exactly which regions YOLO-World identified as flower vs. vase vs.
|
||||
ribbon/bow."""
|
||||
overlay = bgr.copy()
|
||||
for name, boxes in by_class.items():
|
||||
color = _BOX_COLORS.get(name, (170, 170, 170))
|
||||
for box in boxes:
|
||||
x1, y1, x2, y2 = [int(v) for v in box["xyxy"]]
|
||||
cv2.rectangle(overlay, (x1, y1), (x2, y2), color, 3)
|
||||
label = f"{name} {box['confidence']:.2f}"
|
||||
(tw, th), _ = cv2.getTextSize(label, cv2.FONT_HERSHEY_SIMPLEX, 0.7, 2)
|
||||
label_y = max(th + 6, y1)
|
||||
cv2.rectangle(overlay, (x1, label_y - th - 8), (x1 + tw + 6, label_y), color, -1)
|
||||
cv2.putText(overlay, label, (x1 + 3, label_y - 5),
|
||||
cv2.FONT_HERSHEY_SIMPLEX, 0.7, (255, 255, 255), 2, cv2.LINE_AA)
|
||||
return overlay
|
||||
Reference in New Issue
Block a user