Skip to content
21 changes: 19 additions & 2 deletions luxonis_eval/__main__.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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}. "
Expand Down Expand Up @@ -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
)

Expand Down
3 changes: 2 additions & 1 deletion luxonis_eval/utils/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
2 changes: 2 additions & 0 deletions luxonis_eval/visualizers/__init__.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
from .base_visualizer import BaseVisualizer
from .classification import ClassificationVisualizer

__all__ = [
"BaseVisualizer",
"ClassificationVisualizer",
]
51 changes: 49 additions & 2 deletions luxonis_eval/visualizers/base_visualizer.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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()
Loading