diff --git a/luxonis_eval/__main__.py b/luxonis_eval/__main__.py index 01f2558..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}. " @@ -333,7 +347,10 @@ def eval_run( if visualizer: visualizer.visualize( 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/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/__init__.py b/luxonis_eval/visualizers/__init__.py index a5aa25b..b1ac58d 100755 --- a/luxonis_eval/visualizers/__init__.py +++ b/luxonis_eval/visualizers/__init__.py @@ -1,5 +1,7 @@ from .base_visualizer import BaseVisualizer +from .classification import ClassificationVisualizer __all__ = [ "BaseVisualizer", + "ClassificationVisualizer", ] diff --git a/luxonis_eval/visualizers/base_visualizer.py b/luxonis_eval/visualizers/base_visualizer.py index a1dd7e1..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,20 +17,65 @@ 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( 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. + """ ... + + 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 new file mode 100755 index 0000000..369f45c --- /dev/null +++ b/luxonis_eval/visualizers/classification.py @@ -0,0 +1,279 @@ +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) + + 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] + 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 + + 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, + ) + self._draw_text( + vis_frame, + text, + (margin, y_offset + i * line_height), + (0, 255, 0), + font, + font_scale, + thickness, + outline_thickness, + ) + + if target is not None: + 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, + 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, + ) + + self._render( + vis_frame, + f"vis_result_{self.num_visualized:05d}.png", + "Classification Visualization", + ) + 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, + 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, + font_scale: float, + thickness: int, + outline_thickness: int, + ) -> 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 + 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, + 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()