Skip to content

Mmdet Model

sahi.models.mmdet

MMDetection detection model wrapper for SAHI.

Provides integration with OpenMMLab's MMDetection framework for object detection and instance segmentation.

Classes

DetInferencerWrapper

DetInferencerWrapper(
    model: ModelType | str | None = None,
    weights: str | None = None,
    device: str | None = None,
    scope: str | None = "mmdet",
    palette: str = "none",
    image_size: int | None = None,
)

Bases: DetInferencer

Wrapper around MMDetection DetInferencer for custom inference pipeline.

Initialize the DetInferencer wrapper.

Source code in sahi/models/mmdet.py
def __init__(
    self,
    model: ModelType | str | None = None,
    weights: str | None = None,
    device: str | None = None,
    scope: str | None = "mmdet",
    palette: str = "none",
    image_size: int | None = None,
) -> None:
    """Initialize the DetInferencer wrapper."""
    self.image_size = image_size
    super().__init__(model, weights, device, scope, palette)

MmdetDetectionModel

MmdetDetectionModel(
    model_path: str | None = None,
    model: object | None = None,
    config_path: str | None = None,
    device: str | None = None,
    mask_threshold: float = 0.5,
    confidence_threshold: float = 0.3,
    category_mapping: dict | None = None,
    category_remapping: dict | None = None,
    load_at_init: bool = True,
    image_size: int | None = None,
    scope: str = "mmdet",
)

Bases: DetectionModel

MMDetection object detection model.

Wraps MMDetection's DetInferencer for detection and instance segmentation.

Initialize MMDetection detection model.

Source code in sahi/models/mmdet.py
def __init__(
    self,
    model_path: str | None = None,
    model: object | None = None,
    config_path: str | None = None,
    device: str | None = None,
    mask_threshold: float = 0.5,
    confidence_threshold: float = 0.3,
    category_mapping: dict | None = None,
    category_remapping: dict | None = None,
    load_at_init: bool = True,
    image_size: int | None = None,
    scope: str = "mmdet",
) -> None:
    """Initialize MMDetection detection model."""
    self.scope = scope
    self.image_size = image_size
    existing_packages = getattr(self, "required_packages", None) or []
    self.required_packages = [*list(existing_packages), "mmdet", "mmcv", "torch"]
    super().__init__(
        model_path,
        model,
        config_path,
        device,
        mask_threshold,
        confidence_threshold,
        category_mapping,
        category_remapping,
        load_at_init,
        image_size,
    )
Attributes
num_categories property
num_categories: int

Returns number of categories.

has_mask property
has_mask: bool

Returns if model output contains segmentation mask.

Considers both single dataset and ConcatDataset scenarios.

category_names property
category_names: tuple | list

Return category names from model metadata.

Methods:
load_model
load_model() -> None

Detection model is initialized and set to self.model.

Source code in sahi/models/mmdet.py
def load_model(self) -> None:
    """Detection model is initialized and set to self.model."""
    # create model
    model = DetInferencerWrapper(
        self.config_path, self.model_path, device=str(self.device), scope=self.scope, image_size=self.image_size
    )

    self.set_model(model)
set_model
set_model(model: Any, **kwargs: Any) -> None

Sets the underlying MMDetection model.

Parameters:

Name Type Description Default
model Any

Any A MMDetection model

required
**kwargs Any

Any Additional keyword arguments for model setup.

{}
Source code in sahi/models/mmdet.py
def set_model(self, model: Any, **kwargs: Any) -> None:
    """Sets the underlying MMDetection model.

    Args:
        model: Any
            A MMDetection model
        **kwargs: Any
            Additional keyword arguments for model setup.
    """
    # set self.model
    self.model = model

    # set category_mapping
    if not self.category_mapping:
        category_mapping = {str(ind): category_name for ind, category_name in enumerate(self.category_names)}
        self.category_mapping = category_mapping
perform_inference
perform_inference(image: ndarray) -> None

Prediction is performed using self.model and the prediction result is set to self._original_predictions.

Parameters:

Name Type Description Default
image ndarray

np.ndarray A numpy array that contains the image to be predicted. 3 channel image should be in RGB order.

required
Source code in sahi/models/mmdet.py
def perform_inference(self, image: np.ndarray) -> None:
    """Prediction is performed using self.model and the prediction result is set to self._original_predictions.

    Args:
        image: np.ndarray
            A numpy array that contains the image to be predicted. 3 channel image should be in RGB order.
    """
    # Confirm model is loaded
    if self.model is None:
        raise ValueError("Model is not loaded, load it by calling .load_model()")

    # Supports only batch of 1

    # perform inference
    if isinstance(image, np.ndarray):
        # https://github.com/obss/sahi/issues/265
        image = image[:, :, ::-1]
    # compatibility with sahi v0.8.15
    if not isinstance(image, list):
        image_list = [image]
    prediction_result = self.model(image_list)

    self._original_predictions = prediction_result["predictions"]

Functions: