"""Run several detectors over a frame and merge their detections. Used for the multi-model (ADetailer-style) setup: each ticked ``models/yolo//*.pt`` becomes a :class:`~.yolo.YoloDetector` (tagged with its category), and this detector concatenates all their results. Detections keep their own ``label`` (category), so the overlay/table show every model's output together — no cross-model dedup (different categories are meant to coexist). """ from __future__ import annotations from ..video.frame import Frame from .base import Detector from .types import Detection class MultiYoloDetector(Detector): def __init__(self, detectors: list[Detector], nms_iou: float | None = None) -> None: if not detectors: raise ValueError("MultiYoloDetector requires at least one detector") self._detectors = detectors # When set, overlapping detections (across all models, regardless of category) # are merged by greedy IoU NMS — the higher-score box wins. None = keep all. self._nms_iou = nms_iou @property def name(self) -> str: return "Multi(" + " + ".join(d.name for d in self._detectors) + ")" def detect(self, frame: Frame) -> list[Detection]: out: list[Detection] = [] for d in self._detectors: out.extend(d.detect(frame)) if self._nms_iou is not None: out = _nms(out, self._nms_iou) return out def _iou(a: tuple[int, int, int, int], b: tuple[int, int, int, int]) -> float: """Intersection-over-union of two (x, y, w, h) boxes.""" ax, ay, aw, ah = a bx, by, bw, bh = b ix1, iy1 = max(ax, bx), max(ay, by) ix2, iy2 = min(ax + aw, bx + bw), min(ay + ah, by + bh) iw, ih = max(0, ix2 - ix1), max(0, iy2 - iy1) inter = iw * ih if inter == 0: return 0.0 union = aw * ah + bw * bh - inter return inter / union if union > 0 else 0.0 def _nms(dets: list[Detection], iou_thresh: float) -> list[Detection]: """Greedy non-maximum suppression across all detections (category-agnostic). Highest score first; a box is dropped if it overlaps an already-kept box by more than ``iou_thresh``. Used to remove duplicate boxes from overlapping models. """ kept: list[Detection] = [] for d in sorted(dets, key=lambda x: x.score, reverse=True): if all(_iou(d.bbox, k.bbox) <= iou_thresh for k in kept): kept.append(d) return kept