"""Test suite for model functions.""" from pathlib import Path from typing import List import numpy as np import pytest from hypothesis import given, settings from hypothesis import strategies as st from batdetect2 import api from batdetect2.detector import parameters from batdetect2.models.backbones import UNetBackboneConfig from batdetect2.train import load_model_from_checkpoint @settings(deadline=None, max_examples=5) @given(duration=st.floats(min_value=0.1, max_value=2)) @pytest.mark.slow def test_can_import_model_without_pickle(duration: float): # NOTE: remove this test once no other issues are found This is a temporary # test to check that change in model loading did not impact model behaviour # in any way. samplerate = parameters.TARGET_SAMPLERATE_HZ audio = np.random.rand(int(duration * samplerate)) model_without_pickle, model_params_without_pickle = api.load_model( weights_only=True ) model_with_pickle, model_params_with_pickle = api.load_model( weights_only=False ) assert model_params_without_pickle == model_params_with_pickle predictions_without_pickle, _, _ = api.process_audio( audio, model=model_without_pickle, ) predictions_with_pickle, _, _ = api.process_audio( audio, model=model_with_pickle, ) assert predictions_without_pickle == predictions_with_pickle @pytest.mark.slow def test_can_import_model_without_pickle_on_test_data( example_audio_files: List[Path], ): # NOTE: remove this test once no other issues are found This is a temporary # test to check that change in model loading did not impact model behaviour # in any way. model_without_pickle, model_params_without_pickle = api.load_model( weights_only=True ) model_with_pickle, model_params_with_pickle = api.load_model( weights_only=False ) assert model_params_without_pickle == model_params_with_pickle for audio_file in example_audio_files: audio = api.load_audio(str(audio_file)) predictions_without_pickle, _, _ = api.process_audio( audio, model=model_without_pickle, ) predictions_with_pickle, _, _ = api.process_audio( audio, model=model_with_pickle, ) assert predictions_without_pickle == predictions_with_pickle def test_bundled_checkpoint_loads_with_current_config_schema() -> None: """Bundled checkpoints remain loadable after config schema changes.""" model, configs = load_model_from_checkpoint() assert model.class_names assert isinstance(configs.model.architecture, UNetBackboneConfig) assert configs.model.architecture.bottleneck.frequency_aggregation.name