diff --git a/model_api/tests/functional/conftest.py b/model_api/tests/functional/conftest.py index b1eadd93..32bf03e9 100644 --- a/model_api/tests/functional/conftest.py +++ b/model_api/tests/functional/conftest.py @@ -34,6 +34,12 @@ def pytest_addoption(parser): default="", help="directory to store inference result", ) + parser.addoption( + "--only-model-type", + action="store", + default="", + help="select model type to run tests on (and all that inherit from it)", + ) def pytest_configure(config): diff --git a/model_api/tests/functional/test_inference.py b/model_api/tests/functional/test_inference.py index 6e1c4204..3d96369a 100644 --- a/model_api/tests/functional/test_inference.py +++ b/model_api/tests/functional/test_inference.py @@ -6,6 +6,7 @@ import json import operator from pathlib import Path +from typing import Type import cv2 import numpy as np @@ -32,6 +33,7 @@ InstanceSegmentationResult, KeypointDetectionModel, MaskRCNNModel, + Model, Prompt, SAMDecoder, SAMImageEncoder, @@ -145,6 +147,14 @@ def model_data_file(pytestconfig): return pytestconfig.getoption("model_data") +@pytest.fixture(scope="session") +def only_model_class(pytestconfig) -> Type[Model] | None: + model_type = pytestconfig.getoption("only_model_type") + if not model_type: + return None + return Model.get_model_class(model_type) + + def pytest_generate_tests(metafunc): if "model_data" in metafunc.fixturenames: model_data_file = metafunc.config.getoption("model_data") @@ -416,13 +426,18 @@ def assert_contours_match(actual: list[dict], expected: list[dict]) -> None: ) -def test_image_models(data, device, dump, result, model_data, results_dir): # noqa: C901 +def test_image_models(data, device, dump, result, model_data, results_dir, only_model_class): # noqa: C901 name = model_data["name"] + + model_type = MODEL_TYPE_MAPPING[model_data["type"]] + if only_model_class and not issubclass(model_type, only_model_class): + pytest.skip(f"Skipping {name} as it is not a subclass of {only_model_class.__name__}") + if name.endswith((".xml", ".onnx")): name = f"{data}/{name}" for model in create_models( - MODEL_TYPE_MAPPING[model_data["type"]], + model_type, name, data, model_data.get("force_ort", False),