From 4fc1a643999a4abb078d0805eb2adfbede683a85 Mon Sep 17 00:00:00 2001 From: Petros Toupas Date: Thu, 5 Mar 2026 13:04:10 +0200 Subject: [PATCH 1/4] Initial support for the classification visualizer. --- luxonis_eval/__main__.py | 2 + luxonis_eval/visualizers/__init__.py | 5 + luxonis_eval/visualizers/base_visualizer.py | 15 +- luxonis_eval/visualizers/classification.py | 251 ++++++++++++++++++++ 4 files changed, 272 insertions(+), 1 deletion(-) create mode 100755 luxonis_eval/visualizers/classification.py diff --git a/luxonis_eval/__main__.py b/luxonis_eval/__main__.py index 45bf21f..5b2348f 100644 --- a/luxonis_eval/__main__.py +++ b/luxonis_eval/__main__.py @@ -230,7 +230,9 @@ def eval_run( ): task.visualizer.visualize( predictions, + target, infer_engine.vis_frame(), + **metric_ctx, **eval_cfg.visualizer_cfg.params, ) diff --git a/luxonis_eval/visualizers/__init__.py b/luxonis_eval/visualizers/__init__.py index e69de29..8cd3fa3 100755 --- a/luxonis_eval/visualizers/__init__.py +++ b/luxonis_eval/visualizers/__init__.py @@ -0,0 +1,5 @@ +from .classification import ClassificationVisualizer + +__all__ = [ + "ClassificationVisualizer", +] diff --git a/luxonis_eval/visualizers/base_visualizer.py b/luxonis_eval/visualizers/base_visualizer.py index a1dd7e1..cf4187c 100644 --- a/luxonis_eval/visualizers/base_visualizer.py +++ b/luxonis_eval/visualizers/base_visualizer.py @@ -27,8 +27,21 @@ def __init__(self, **kwargs: Any) -> None: def visualize( self, predictions: Any, + target: Any, vis_frame: np.ndarray, **kwargs: Any, ) -> None: - """Visualize the evaluation results.""" + """Visualize the evaluation results. + + Parameters + ---------- + predictions : Any + Model predictions. + target : Any + Ground truth target. + vis_frame : np.ndarray + Frame to visualize on. + **kwargs : Any + Additional visualization options. + """ ... diff --git a/luxonis_eval/visualizers/classification.py b/luxonis_eval/visualizers/classification.py new file mode 100755 index 0000000..94c5f1d --- /dev/null +++ b/luxonis_eval/visualizers/classification.py @@ -0,0 +1,251 @@ +from typing import Any + +import cv2 +import numpy as np +from depthai_nodes import Classifications + +from luxonis_eval.visualizers.base_visualizer import BaseVisualizer + + +class ClassificationVisualizer(BaseVisualizer): + """Visualizer for classification tasks.""" + + def __init__( + self, *, max_visualizations: int | None = None, **kwargs: Any + ) -> None: + """Initialize the classification visualizer. + + Parameters + ---------- + max_visualizations : int | None, optional + Maximum number of visualizations to display. If None, visualizes all samples. + **kwargs : Any + Additional visualization options. + """ + self.max_visualizations = max_visualizations + self.num_visualized = 0 + super().__init__(**kwargs) + + @staticmethod + def _fit_text( + text: str, + font: int, + font_scale: float, + thickness: int, + max_text_w: int, + ) -> str: + """Truncate text to fit within the image width. + + Parameters + ---------- + text : str + Text to truncate. + font : int + Font to use. + font_scale : float + Font scale. + thickness : int + Text thickness. + max_text_w : int + Maximum width of the text. + + Returns + ------- + str + Truncated text. + """ + (tw, _), _ = cv2.getTextSize(text, font, font_scale, thickness) + if tw <= max_text_w: + return text + while len(text) > 1: + text = text[:-1] + (tw, _), _ = cv2.getTextSize( + text + "...", font, font_scale, thickness + ) + if tw <= max_text_w: + return text + "..." + return text + + @staticmethod + def _draw_text( + frame: np.ndarray, + text: str, + pos: tuple[int, int], + color: tuple[int, int, int], + font: int = cv2.FONT_HERSHEY_SIMPLEX, + font_scale: float = 0.5, + thickness: int = 1, + outline_thickness: int = 3, + ) -> None: + """Draw text with a dark outline for readability. + + Parameters + ---------- + frame : np.ndarray + Image to draw on. + text : str + Text to draw. + pos : tuple[int, int] + Bottom-left corner of the text. + color : tuple[int, int, int] + Text color in BGR format. + font : int, optional + Font to use, by default cv2.FONT_HERSHEY_SIMPLEX + font_scale : float, optional + Font scale, by default 0.5 + thickness : int, optional + Text thickness, by default 1 + outline_thickness : int, optional + Thickness of the text outline, by default 3 + """ + cv2.putText( + frame, + text, + pos, + font, + font_scale, + (0, 0, 0), + outline_thickness, + cv2.LINE_AA, + ) + cv2.putText( + frame, + text, + pos, + font, + font_scale, + color, + thickness, + cv2.LINE_AA, + ) + + @staticmethod + def _short_name(name: str) -> str: + """Keep only the first comma-separated name. + + Parameters + ---------- + name : str + Original name. + Returns + ------- + str + Shortened name. + """ + return name.split(",", maxsplit=1)[0].strip() + + def visualize( + self, + predictions: Classifications, + target: Any, + vis_frame: np.ndarray, + *, + resize_ratio: float | None = None, + **kwargs: Any, + ) -> None: + """Visualize the classification results. + + Parameters + ---------- + predictions : Classifications + Model predictions. + target : Any + Ground truth target. + vis_frame : np.ndarray + Frame to visualize on. + resize_ratio : float | None, optional + Ratio to resize the visualization frame by. + **kwargs : Any + Additional visualization options. + """ + if ( + self.max_visualizations is not None + and self.num_visualized >= self.max_visualizations + ): + return + + if resize_ratio is not None: + h, w = vis_frame.shape[:2] + new_w = int(w * resize_ratio) + new_h = int(h * resize_ratio) + vis_frame = cv2.resize(vis_frame, (new_w, new_h)) + + top_k = kwargs.get("top_k", 5) + font = cv2.FONT_HERSHEY_SIMPLEX + img_w = vis_frame.shape[1] + font_scale = max(0.3, img_w / 1000) + thickness = max(1, int(img_w / 500)) + line_height = int(font_scale * 40) + y_offset = line_height + 5 + margin = 10 + max_text_w = img_w - 2 * margin + + outline_thickness = thickness + 2 + + pred_classes = predictions.classes[:top_k] + pred_scores = predictions.scores[:top_k] + + for i, (cls_name, score) in enumerate( + zip(pred_classes, pred_scores, strict=True) + ): + text = self._fit_text( + f"{self._short_name(cls_name)}: {score:.2%}", + font, + font_scale, + thickness, + max_text_w, + ) + y = y_offset + i * line_height + self._draw_text( + vis_frame, + text, + (margin, y), + (0, 255, 0), + font, + font_scale, + thickness, + outline_thickness, + ) + + if target is not None: + class_index_map = kwargs.get("class_index_map") + class_map = kwargs.get("class_map", {}) + inv_class_map = {v: k for k, v in class_map.items()} + + cls_target = target.get("/classification") + if cls_target is not None: + tgt = np.asarray(cls_target) + target_idx = ( + int(np.argmax(tgt)) + if tgt.ndim > 0 and tgt.size > 1 + else int(tgt) + ) + if class_index_map is not None: + target_idx = int(class_index_map[target_idx]) + gt_label = inv_class_map.get(target_idx, str(target_idx)) + else: + gt_label = str(target) + + gt_text = self._fit_text( + f"GT: {self._short_name(gt_label)}", + font, + font_scale, + thickness, + max_text_w, + ) + gt_y = y_offset + len(pred_classes) * line_height + 10 + self._draw_text( + vis_frame, + gt_text, + (margin, gt_y), + (0, 0, 255), + font, + font_scale, + thickness, + outline_thickness, + ) + + cv2.imshow("Classification Visualization", vis_frame) + cv2.waitKey(0) + cv2.destroyAllWindows() + self.num_visualized += 1 From 211a95ef67ff4fb8436294208c794e09ee12b09c Mon Sep 17 00:00:00 2001 From: Petros Toupas Date: Fri, 17 Apr 2026 12:52:53 +0300 Subject: [PATCH 2/4] Fix minor visualization issues on ClassificationVisualizer --- luxonis_eval/__main__.py | 2 + luxonis_eval/visualizers/classification.py | 230 +++++++++++---------- 2 files changed, 118 insertions(+), 114 deletions(-) diff --git a/luxonis_eval/__main__.py b/luxonis_eval/__main__.py index 7985b3a..a83912f 100644 --- a/luxonis_eval/__main__.py +++ b/luxonis_eval/__main__.py @@ -335,6 +335,8 @@ def eval_run( predictions, target, infer_engine.vis_frame(), + class_map=class_map, + class_index_map=class_index_map, **eval_cfg.visualizer.params, # type: ignore ) diff --git a/luxonis_eval/visualizers/classification.py b/luxonis_eval/visualizers/classification.py index 94c5f1d..3961a28 100755 --- a/luxonis_eval/visualizers/classification.py +++ b/luxonis_eval/visualizers/classification.py @@ -26,114 +26,6 @@ def __init__( self.num_visualized = 0 super().__init__(**kwargs) - @staticmethod - def _fit_text( - text: str, - font: int, - font_scale: float, - thickness: int, - max_text_w: int, - ) -> str: - """Truncate text to fit within the image width. - - Parameters - ---------- - text : str - Text to truncate. - font : int - Font to use. - font_scale : float - Font scale. - thickness : int - Text thickness. - max_text_w : int - Maximum width of the text. - - Returns - ------- - str - Truncated text. - """ - (tw, _), _ = cv2.getTextSize(text, font, font_scale, thickness) - if tw <= max_text_w: - return text - while len(text) > 1: - text = text[:-1] - (tw, _), _ = cv2.getTextSize( - text + "...", font, font_scale, thickness - ) - if tw <= max_text_w: - return text + "..." - return text - - @staticmethod - def _draw_text( - frame: np.ndarray, - text: str, - pos: tuple[int, int], - color: tuple[int, int, int], - font: int = cv2.FONT_HERSHEY_SIMPLEX, - font_scale: float = 0.5, - thickness: int = 1, - outline_thickness: int = 3, - ) -> None: - """Draw text with a dark outline for readability. - - Parameters - ---------- - frame : np.ndarray - Image to draw on. - text : str - Text to draw. - pos : tuple[int, int] - Bottom-left corner of the text. - color : tuple[int, int, int] - Text color in BGR format. - font : int, optional - Font to use, by default cv2.FONT_HERSHEY_SIMPLEX - font_scale : float, optional - Font scale, by default 0.5 - thickness : int, optional - Text thickness, by default 1 - outline_thickness : int, optional - Thickness of the text outline, by default 3 - """ - cv2.putText( - frame, - text, - pos, - font, - font_scale, - (0, 0, 0), - outline_thickness, - cv2.LINE_AA, - ) - cv2.putText( - frame, - text, - pos, - font, - font_scale, - color, - thickness, - cv2.LINE_AA, - ) - - @staticmethod - def _short_name(name: str) -> str: - """Keep only the first comma-separated name. - - Parameters - ---------- - name : str - Original name. - Returns - ------- - str - Shortened name. - """ - return name.split(",", maxsplit=1)[0].strip() - def visualize( self, predictions: Classifications, @@ -210,19 +102,21 @@ def visualize( if target is not None: class_index_map = kwargs.get("class_index_map") class_map = kwargs.get("class_map", {}) - inv_class_map = {v: k for k, v in class_map.items()} - cls_target = target.get("/classification") + cls_target = target.get("/classification") if isinstance(target, dict) else None if cls_target is not None: tgt = np.asarray(cls_target) - target_idx = ( + ldf_idx = ( int(np.argmax(tgt)) if tgt.ndim > 0 and tgt.size > 1 else int(tgt) ) - if class_index_map is not None: - target_idx = int(class_index_map[target_idx]) - gt_label = inv_class_map.get(target_idx, str(target_idx)) + target_idx = ( + int(class_index_map[ldf_idx]) + if class_index_map is not None + else ldf_idx + ) + gt_label = class_map.get(target_idx, str(target_idx)) else: gt_label = str(target) @@ -249,3 +143,111 @@ def visualize( cv2.waitKey(0) cv2.destroyAllWindows() self.num_visualized += 1 + + @staticmethod + def _fit_text( + text: str, + font: int, + font_scale: float, + thickness: int, + max_text_w: int, + ) -> str: + """Truncate text to fit within the image width. + + Parameters + ---------- + text : str + Text to truncate. + font : int + Font to use. + font_scale : float + Font scale. + thickness : int + Text thickness. + max_text_w : int + Maximum width of the text. + + Returns + ------- + str + Truncated text. + """ + (tw, _), _ = cv2.getTextSize(text, font, font_scale, thickness) + if tw <= max_text_w: + return text + while len(text) > 1: + text = text[:-1] + (tw, _), _ = cv2.getTextSize( + text + "...", font, font_scale, thickness + ) + if tw <= max_text_w: + return text + "..." + return text + + @staticmethod + def _draw_text( + frame: np.ndarray, + text: str, + pos: tuple[int, int], + color: tuple[int, int, int], + font: int = cv2.FONT_HERSHEY_SIMPLEX, + font_scale: float = 0.5, + thickness: int = 1, + outline_thickness: int = 3, + ) -> None: + """Draw text with a dark outline for readability. + + Parameters + ---------- + frame : np.ndarray + Image to draw on. + text : str + Text to draw. + pos : tuple[int, int] + Bottom-left corner of the text. + color : tuple[int, int, int] + Text color in BGR format. + font : int, optional + Font to use, by default cv2.FONT_HERSHEY_SIMPLEX + font_scale : float, optional + Font scale, by default 0.5 + thickness : int, optional + Text thickness, by default 1 + outline_thickness : int, optional + Thickness of the text outline, by default 3 + """ + cv2.putText( + frame, + text, + pos, + font, + font_scale, + (0, 0, 0), + outline_thickness, + cv2.LINE_AA, + ) + cv2.putText( + frame, + text, + pos, + font, + font_scale, + color, + thickness, + cv2.LINE_AA, + ) + + @staticmethod + def _short_name(name: str) -> str: + """Keep only the first comma-separated name. + + Parameters + ---------- + name : str + Original name. + Returns + ------- + str + Shortened name. + """ + return name.split(",", maxsplit=1)[0].strip() From 178b0ae6b5eb37bd362db5815698946a1db0c022 Mon Sep 17 00:00:00 2001 From: Petros Toupas Date: Fri, 17 Apr 2026 13:25:46 +0300 Subject: [PATCH 3/4] Tidy up ClassificationVisualizer --- luxonis_eval/visualizers/classification.py | 102 +++++++++++++-------- 1 file changed, 63 insertions(+), 39 deletions(-) diff --git a/luxonis_eval/visualizers/classification.py b/luxonis_eval/visualizers/classification.py index 3961a28..925798c 100755 --- a/luxonis_eval/visualizers/classification.py +++ b/luxonis_eval/visualizers/classification.py @@ -58,22 +58,24 @@ def visualize( if resize_ratio is not None: h, w = vis_frame.shape[:2] - new_w = int(w * resize_ratio) - new_h = int(h * resize_ratio) - vis_frame = cv2.resize(vis_frame, (new_w, new_h)) + vis_frame = cv2.resize( + vis_frame, (int(w * resize_ratio), int(h * resize_ratio)) + ) top_k = kwargs.get("top_k", 5) + class_map = kwargs.get("class_map", {}) + class_index_map = kwargs.get("class_index_map") + font = cv2.FONT_HERSHEY_SIMPLEX img_w = vis_frame.shape[1] font_scale = max(0.3, img_w / 1000) thickness = max(1, int(img_w / 500)) + outline_thickness = thickness + 2 line_height = int(font_scale * 40) y_offset = line_height + 5 margin = 10 max_text_w = img_w - 2 * margin - outline_thickness = thickness + 2 - pred_classes = predictions.classes[:top_k] pred_scores = predictions.scores[:top_k] @@ -87,11 +89,10 @@ def visualize( thickness, max_text_w, ) - y = y_offset + i * line_height self._draw_text( vis_frame, text, - (margin, y), + (margin, y_offset + i * line_height), (0, 255, 0), font, font_scale, @@ -100,26 +101,9 @@ def visualize( ) if target is not None: - class_index_map = kwargs.get("class_index_map") - class_map = kwargs.get("class_map", {}) - - cls_target = target.get("/classification") if isinstance(target, dict) else None - if cls_target is not None: - tgt = np.asarray(cls_target) - ldf_idx = ( - int(np.argmax(tgt)) - if tgt.ndim > 0 and tgt.size > 1 - else int(tgt) - ) - target_idx = ( - int(class_index_map[ldf_idx]) - if class_index_map is not None - else ldf_idx - ) - gt_label = class_map.get(target_idx, str(target_idx)) - else: - gt_label = str(target) - + gt_label = self._resolve_gt_label( + target, class_map=class_map, class_index_map=class_index_map + ) gt_text = self._fit_text( f"GT: {self._short_name(gt_label)}", font, @@ -144,6 +128,45 @@ def visualize( cv2.destroyAllWindows() self.num_visualized += 1 + @staticmethod + def _resolve_gt_label( + target: Any, + *, + class_map: dict[int, str], + class_index_map: dict[int, int] | None, + ) -> str: + """Resolve the ground-truth label string from a raw target. + + Parameters + ---------- + target : Any + Ground truth target, either a dict with a ``/classification`` key or a raw value. + class_map : dict[int, str] + Mapping from model class index to class name. + class_index_map : dict[int, int] | None + Optional mapping from dataset index to model class index. + + Returns + ------- + str + Human-readable ground-truth label. + """ + cls_target = ( + target.get("/classification") if isinstance(target, dict) else None + ) + if cls_target is None: + return str(target) + tgt = np.asarray(cls_target) + ldf_idx = ( + int(np.argmax(tgt)) if tgt.ndim > 0 and tgt.size > 1 else int(tgt) + ) + target_idx = ( + int(class_index_map[ldf_idx]) + if class_index_map is not None + else ldf_idx + ) + return class_map.get(target_idx, str(target_idx)) + @staticmethod def _fit_text( text: str, @@ -190,10 +213,10 @@ def _draw_text( text: str, pos: tuple[int, int], color: tuple[int, int, int], - font: int = cv2.FONT_HERSHEY_SIMPLEX, - font_scale: float = 0.5, - thickness: int = 1, - outline_thickness: int = 3, + font: int, + font_scale: float, + thickness: int, + outline_thickness: int, ) -> None: """Draw text with a dark outline for readability. @@ -207,14 +230,14 @@ def _draw_text( Bottom-left corner of the text. color : tuple[int, int, int] Text color in BGR format. - font : int, optional - Font to use, by default cv2.FONT_HERSHEY_SIMPLEX - font_scale : float, optional - Font scale, by default 0.5 - thickness : int, optional - Text thickness, by default 1 - outline_thickness : int, optional - Thickness of the text outline, by default 3 + font : int + OpenCV font identifier. + font_scale : float + Font scale. + thickness : int + Text thickness. + outline_thickness : int + Thickness of the dark outline drawn behind the text. """ cv2.putText( frame, @@ -245,6 +268,7 @@ def _short_name(name: str) -> str: ---------- name : str Original name. + Returns ------- str From 62cc3ffd02b3de1db04fc34524d4caf70b9a411c Mon Sep 17 00:00:00 2001 From: Petros Toupas Date: Fri, 17 Apr 2026 14:09:39 +0300 Subject: [PATCH 4/4] Update VisualizerConfig and add a new option to choose between saving and displaying interactively. --- luxonis_eval/__main__.py | 18 +++++++++-- luxonis_eval/utils/config.py | 3 +- luxonis_eval/visualizers/base_visualizer.py | 36 ++++++++++++++++++++- luxonis_eval/visualizers/classification.py | 8 +++-- 4 files changed, 58 insertions(+), 7 deletions(-) diff --git a/luxonis_eval/__main__.py b/luxonis_eval/__main__.py index a83912f..85e6c18 100644 --- a/luxonis_eval/__main__.py +++ b/luxonis_eval/__main__.py @@ -1,6 +1,7 @@ import time import types from importlib.metadata import version +from pathlib import Path from typing import Literal import numpy as np @@ -195,17 +196,30 @@ def eval_setup( # ------------------------------------------------------------------------- # Visualizer initialization # ------------------------------------------------------------------------- - if eval_cfg.visualizer and eval_cfg.visualizer.visualize: + if eval_cfg.visualizer: try: + vis_save_dir = ( + Path(eval_cfg.visualizer.save_dir) / get_model_name(eval_cfg.engine.model_path) + if eval_cfg.visualizer.mode == "save" + else None + ) visualizer = from_registry( VISUALIZERS_REGISTRY, eval_cfg.visualizer.name, + save_dir=vis_save_dir, **eval_cfg.visualizer.params, ) assert isinstance(visualizer, BaseVisualizer), ( f"{eval_cfg.visualizer.name} visualizer must be an instance of BaseVisualizer." ) - logger.info(f"{eval_cfg.visualizer.name} visualizer initialized.") + if vis_save_dir is not None: + logger.info( + f"{eval_cfg.visualizer.name} visualizer initialized. Results will be saved to '{vis_save_dir}'." + ) + else: + logger.info( + f"{eval_cfg.visualizer.name} visualizer initialized. Results will be displayed interactively." + ) except KeyError as e: raise ValueError( f"Unknown visualizer: {eval_cfg.visualizer.name}. " diff --git a/luxonis_eval/utils/config.py b/luxonis_eval/utils/config.py index 1c59235..a744420 100644 --- a/luxonis_eval/utils/config.py +++ b/luxonis_eval/utils/config.py @@ -98,7 +98,8 @@ class MetricsConfig(BaseModelExtraForbid): class VisualizerConfig(ConfigItem): - visualize: bool = True + mode: Literal["display", "save"] = "display" + save_dir: str = "visualizations" @field_validator("name", mode="after") def validate_name(cls, v: str) -> str: diff --git a/luxonis_eval/visualizers/base_visualizer.py b/luxonis_eval/visualizers/base_visualizer.py index cf4187c..fdaaf59 100644 --- a/luxonis_eval/visualizers/base_visualizer.py +++ b/luxonis_eval/visualizers/base_visualizer.py @@ -1,6 +1,8 @@ from abc import ABC, abstractmethod +from pathlib import Path from typing import Any +import cv2 import numpy as np from luxonis_ml.utils.registry import AutoRegisterMeta @@ -15,13 +17,19 @@ class BaseVisualizer( ): """Base class for evaluation visualizers.""" - def __init__(self, **kwargs: Any) -> None: + def __init__(self, *, save_dir: Path | None = None, **kwargs: Any) -> None: """Initialize the visualizer. + Parameters ---------- + save_dir : Path | None, optional + Directory to save visualization frames to. **kwargs : Any Visualizer basic configuration. """ + self.save_dir = save_dir + if save_dir is not None: + save_dir.mkdir(parents=True, exist_ok=True) @abstractmethod def visualize( @@ -45,3 +53,29 @@ def visualize( Additional visualization options. """ ... + + def _render( + self, + frame: np.ndarray, + filename: str, + window_title: str = "Visualization", + ) -> None: + """Display a frame interactively or save it to disk. + + Parameters + ---------- + frame : np.ndarray + Image to display or save. + filename : str + Output filename used when saving to disk. + window_title : str, optional + Window title used when displaying interactively. + """ + if self.save_dir is not None: + if frame.dtype != np.uint8: + frame = np.clip(frame * 255.0, 0, 255).astype(np.uint8) + cv2.imwrite(str(self.save_dir / filename), frame) + else: + cv2.imshow(window_title, frame) + cv2.waitKey(0) + cv2.destroyAllWindows() diff --git a/luxonis_eval/visualizers/classification.py b/luxonis_eval/visualizers/classification.py index 925798c..369f45c 100755 --- a/luxonis_eval/visualizers/classification.py +++ b/luxonis_eval/visualizers/classification.py @@ -123,9 +123,11 @@ def visualize( outline_thickness, ) - cv2.imshow("Classification Visualization", vis_frame) - cv2.waitKey(0) - cv2.destroyAllWindows() + self._render( + vis_frame, + f"vis_result_{self.num_visualized:05d}.png", + "Classification Visualization", + ) self.num_visualized += 1 @staticmethod