test: cover compile idempotency paths

This commit is contained in:
mbsantiago 2026-08-08 12:04:19 +01:00
parent 8166db8f9c
commit bad5d4f4fc
4 changed files with 72 additions and 15 deletions

View File

@ -1,4 +1,5 @@
import uuid
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable, List, Optional, cast
from uuid import uuid4
@ -32,6 +33,12 @@ from batdetect2.train.lightning import build_training_module
from batdetect2.train.types import ClipLabeller
@dataclass
class DetectorCompileRecorder:
compile_count: int = 0
call_count: int = 0
@pytest.fixture
def example_data_dir() -> Path:
pkg_dir = Path(__file__).parent.parent
@ -368,23 +375,25 @@ def sample_audio_loader() -> AudioLoader:
@pytest.fixture
def record_compiled_detector_calls(
def record_detector_compilation(
monkeypatch: pytest.MonkeyPatch,
) -> Callable[[ModelProtocol], list[None]]:
def factory(model: ModelProtocol) -> list[None]:
compiled_calls: list[None] = []
) -> Callable[[ModelProtocol], DetectorCompileRecorder]:
def factory(model: ModelProtocol) -> DetectorCompileRecorder:
recorder = DetectorCompileRecorder()
detector = cast(Any, model.detector)
original_call_impl = detector._call_impl
def compile_detector() -> None:
recorder.compile_count += 1
def compiled_call(*args, **kwargs):
compiled_calls.append(None)
recorder.call_count += 1
return original_call_impl(*args, **kwargs)
detector._compiled_call_impl = compiled_call
monkeypatch.setattr(detector, "compile", compile_detector)
return compiled_calls
return recorder
return factory

View File

@ -156,9 +156,9 @@ def test_process_spectrogram_rejects_batched_input(
def test_user_can_compile_api_detector(
api_v2: BatDetect2API,
example_audio_files: list[Path],
record_compiled_detector_calls,
record_detector_compilation,
) -> None:
compiled_calls = record_compiled_detector_calls(api_v2.model)
recorder = record_detector_compilation(api_v2.model)
audio = api_v2.load_audio(example_audio_files[0])
spec = api_v2.generate_spectrogram(audio)
@ -166,7 +166,8 @@ def test_user_can_compile_api_detector(
api_v2.compile()
api_v2.process_spectrogram(spec)
assert len(compiled_calls) == 1
assert recorder.compile_count == 1
assert recorder.call_count == 1
def test_api_from_config_compiles_detector_when_requested(
@ -187,6 +188,26 @@ def test_api_from_config_compiles_detector_when_requested(
assert compiled_models == [api.model]
def test_api_from_checkpoint_compiles_detector_when_requested(
tiny_checkpoint_path: Path,
monkeypatch: pytest.MonkeyPatch,
) -> None:
compiled_models = []
def compile_model(model):
compiled_models.append(model)
return model
monkeypatch.setattr("batdetect2.models.compile_model", compile_model)
api = BatDetect2API.from_checkpoint(
tiny_checkpoint_path,
compile_model=True,
)
assert compiled_models == [api.model]
def test_user_can_read_top_class_and_other_class_scores(
api_v2: BatDetect2API,
example_audio_files: list[Path],

View File

@ -59,10 +59,35 @@ def test_run_batch_inference_matches_single_clip_inference(
def test_run_batch_inference_compiles_detector_when_config_requests_compile(
example_annotations: list[data.ClipAnnotation],
record_compiled_detector_calls,
record_detector_compilation,
) -> None:
api = BatDetect2API.from_config()
compiled_calls = record_compiled_detector_calls(api.model)
recorder = record_detector_compilation(api.model)
predictions = run_batch_inference(
api.model,
[example_annotations[0].clip],
targets=api.targets,
roi_mapper=api.roi_mapper,
audio_loader=api.audio_loader,
preprocessor=api.preprocessor,
output_transform=api.output_transform,
inference_config=InferenceConfig(compile_model=True),
batch_size=1,
num_workers=0,
)
assert predictions
assert recorder.compile_count == 1
assert recorder.call_count > 0
def test_run_batch_inference_does_not_recompile_compiled_detector(
example_annotations: list[data.ClipAnnotation],
record_detector_compilation,
) -> None:
api = BatDetect2API.from_config()
recorder = record_detector_compilation(api.model)
api.compile()
predictions = run_batch_inference(
@ -79,4 +104,5 @@ def test_run_batch_inference_compiles_detector_when_config_requests_compile(
)
assert predictions
assert compiled_calls
assert recorder.compile_count == 1
assert recorder.call_count > 0

View File

@ -312,7 +312,7 @@ def test_train_smoke_produces_loadable_checkpoint(
def test_run_train_compiles_detector_when_train_config_requests_compile(
tmp_path: Path,
example_annotations: list[data.ClipAnnotation],
record_compiled_detector_calls,
record_detector_compilation,
) -> None:
targets_config = TargetConfig()
targets = build_targets(targets_config)
@ -324,7 +324,7 @@ def test_run_train_compiles_detector_when_train_config_requests_compile(
)
train_config = build_fast_train_config()
train_config.compile_model = True
compiled_calls = record_compiled_detector_calls(model)
recorder = record_detector_compilation(model)
module = run_train(
train_annotations=example_annotations[:1],
@ -345,7 +345,8 @@ def test_run_train_compiles_detector_when_train_config_requests_compile(
assert (
getattr(module.model.detector, "_compiled_call_impl", None) is not None
)
assert compiled_calls
assert recorder.compile_count == 1
assert recorder.call_count > 0
@pytest.mark.slow