feat: add runtime model compilation options

This commit is contained in:
mbsantiago 2026-08-08 12:02:00 +01:00
parent 9f22fd097f
commit 8166db8f9c
11 changed files with 229 additions and 33 deletions

View File

@ -136,7 +136,7 @@ clean: clean-build clean-pyc clean-test clean-docs
# Train on example data. # Train on example data.
example-train OPTIONS="": example-train OPTIONS="":
uv run batdetect2 train \ uv run batdetect2 -v train \
--val-dataset example_data/dataset.yaml \ --val-dataset example_data/dataset.yaml \
--base-dir . \ --base-dir . \
--targets example_data/targets.yaml \ --targets example_data/targets.yaml \

View File

@ -1,5 +1,4 @@
from __future__ import annotations from __future__ import annotations
from lightning.pytorch.loggers import Logger
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Literal from typing import TYPE_CHECKING, Literal
@ -9,6 +8,7 @@ if TYPE_CHECKING:
import numpy as np import numpy as np
import torch import torch
from lightning.pytorch.loggers import Logger
from soundevent import data from soundevent import data
from batdetect2.audio import AudioConfig, AudioLoader from batdetect2.audio import AudioConfig, AudioLoader
@ -153,6 +153,19 @@ class BatDetect2API:
self.model.eval() self.model.eval()
def compile(self) -> "BatDetect2API":
"""Compile the detector path used by inference.
Returns
-------
BatDetect2API
This API instance with the detector compiled.
"""
from batdetect2.models import compile_model
compile_model(self.model)
return self
def load_annotations( def load_annotations(
self, self,
path: data.PathLike, path: data.PathLike,
@ -982,6 +995,7 @@ class BatDetect2API:
inference_config: InferenceConfig | None = None, inference_config: InferenceConfig | None = None,
outputs_config: OutputsConfig | None = None, outputs_config: OutputsConfig | None = None,
logging_config: AppLoggingConfig | None = None, logging_config: AppLoggingConfig | None = None,
compile_model: bool = False,
) -> "BatDetect2API": ) -> "BatDetect2API":
"""Build an API instance from config objects. """Build an API instance from config objects.
@ -1007,6 +1021,8 @@ class BatDetect2API:
Output config. If omitted, the default outputs config is used. Output config. If omitted, the default outputs config is used.
logging_config : AppLoggingConfig | None, optional logging_config : AppLoggingConfig | None, optional
Logging config. If omitted, the default logging config is used. Logging config. If omitted, the default logging config is used.
compile_model : bool, optional
If ``True``, compile the detector path after building the API.
Returns Returns
------- -------
@ -1089,7 +1105,7 @@ class BatDetect2API:
), ),
) )
return cls( api = cls(
model_config=model_config, model_config=model_config,
audio_config=audio_config, audio_config=audio_config,
train_config=train_config, train_config=train_config,
@ -1108,6 +1124,11 @@ class BatDetect2API:
output_transform=output_transform, output_transform=output_transform,
) )
if compile_model:
api.compile()
return api
@classmethod @classmethod
def from_checkpoint( def from_checkpoint(
cls, cls,
@ -1118,6 +1139,7 @@ class BatDetect2API:
inference_config: InferenceConfig | None = None, inference_config: InferenceConfig | None = None,
outputs_config: OutputsConfig | None = None, outputs_config: OutputsConfig | None = None,
logging_config: AppLoggingConfig | None = None, logging_config: AppLoggingConfig | None = None,
compile_model: bool = False,
) -> "BatDetect2API": ) -> "BatDetect2API":
"""Build an API instance from a saved checkpoint. """Build an API instance from a saved checkpoint.
@ -1138,6 +1160,8 @@ class BatDetect2API:
Output config override. Output config override.
logging_config : AppLoggingConfig | None, optional logging_config : AppLoggingConfig | None, optional
Logging config override. Logging config override.
compile_model : bool, optional
If ``True``, compile the detector path after building the API.
Returns Returns
------- -------
@ -1221,7 +1245,7 @@ class BatDetect2API:
transform=output_transform, transform=output_transform,
) )
return cls( api = cls(
model_config=model_config, model_config=model_config,
audio_config=audio_config, audio_config=audio_config,
train_config=train_config, train_config=train_config,
@ -1240,6 +1264,11 @@ class BatDetect2API:
output_transform=output_transform, output_transform=output_transform,
) )
if compile_model:
api.compile()
return api
def _set_trainable_parameters( def _set_trainable_parameters(
self, self,
trainable: Literal["all", "heads", "classifier_head", "size_head"], trainable: Literal["all", "heads", "classifier_head", "size_head"],

View File

@ -10,6 +10,7 @@ from batdetect2.inference.clips import get_clips_from_files
from batdetect2.inference.config import InferenceConfig from batdetect2.inference.config import InferenceConfig
from batdetect2.inference.dataset import build_inference_loader from batdetect2.inference.dataset import build_inference_loader
from batdetect2.inference.lightning import InferenceModule from batdetect2.inference.lightning import InferenceModule
from batdetect2.models import compile_model
from batdetect2.models.types import ModelProtocol from batdetect2.models.types import ModelProtocol
from batdetect2.outputs import ( from batdetect2.outputs import (
OutputsConfig, OutputsConfig,
@ -71,6 +72,9 @@ def run_batch_inference(
batch_size=batch_size, batch_size=batch_size,
) )
if inference_config.compile_model:
compile_model(model)
module = InferenceModule( module = InferenceModule(
model, model,
output_transform=output_transform, output_transform=output_transform,

View File

@ -15,6 +15,7 @@ class ClipingConfig(BaseConfig):
class InferenceConfig(BaseConfig): class InferenceConfig(BaseConfig):
compile_model: bool = False
loader: InferenceLoaderConfig = Field( loader: InferenceLoaderConfig = Field(
default_factory=InferenceLoaderConfig default_factory=InferenceLoaderConfig
) )

View File

@ -100,6 +100,7 @@ __all__ = [
"ModelConfig", "ModelConfig",
"build_model", "build_model",
"build_model_with_new_targets", "build_model_with_new_targets",
"compile_model",
] ]
@ -112,9 +113,6 @@ class ModelConfig(BaseConfig):
Attributes Attributes
---------- ----------
compile : bool
If ``True``, compile the model before training. Defaults to
``False``.
samplerate : int samplerate : int
Expected input audio sample rate in Hz. Audio must be resampled Expected input audio sample rate in Hz. Audio must be resampled
to this rate before being passed to the model. Defaults to to this rate before being passed to the model. Defaults to
@ -132,7 +130,6 @@ class ModelConfig(BaseConfig):
``PostprocessConfig()``. ``PostprocessConfig()``.
""" """
compile: bool = False
samplerate: int = Field(default=TARGET_SAMPLERATE_HZ, gt=0) samplerate: int = Field(default=TARGET_SAMPLERATE_HZ, gt=0)
architecture: BackboneConfig = Field(default_factory=UNetBackboneConfig) architecture: BackboneConfig = Field(default_factory=UNetBackboneConfig)
preprocess: PreprocessingConfig = Field( preprocess: PreprocessingConfig = Field(
@ -323,3 +320,15 @@ def build_model_with_new_targets(
dimension_names=roi_mapper.dimension_names, dimension_names=roi_mapper.dimension_names,
config=model.get_config(), config=model.get_config(),
) )
def compile_model(model: ModelProtocol) -> ModelProtocol:
"""Compile the detector path used by training and inference."""
if not isinstance(model.detector, torch.nn.Module):
raise TypeError("Detector must be a torch.nn.Module to compile.")
if getattr(model.detector, "_compiled_call_impl", None) is not None:
return model
model.detector.compile()
return model

View File

@ -41,6 +41,7 @@ class PLTrainerConfig(BaseConfig):
class TrainingConfig(BaseConfig): class TrainingConfig(BaseConfig):
compile_model: bool = False
precision: Literal["medium", "high"] | None = None precision: Literal["medium", "high"] | None = None
train_loader: TrainLoaderConfig = Field(default_factory=TrainLoaderConfig) train_loader: TrainLoaderConfig = Field(default_factory=TrainLoaderConfig)
val_loader: ValLoaderConfig = Field(default_factory=ValLoaderConfig) val_loader: ValLoaderConfig = Field(default_factory=ValLoaderConfig)

View File

@ -16,7 +16,7 @@ from batdetect2.logging import (
LoggingCallback, LoggingCallback,
build_logger, build_logger,
) )
from batdetect2.models import ModelConfig, build_model from batdetect2.models import ModelConfig, build_model, compile_model
from batdetect2.models.types import ModelProtocol from batdetect2.models.types import ModelProtocol
from batdetect2.preprocess import PreprocessorProtocol, build_preprocessor from batdetect2.preprocess import PreprocessorProtocol, build_preprocessor
from batdetect2.targets import ( from batdetect2.targets import (
@ -208,17 +208,17 @@ def run_train(
run_name=run_name, run_name=run_name,
) )
if model_config.compile:
logger.info("Compiling model...")
module.compile()
if train_config.precision is not None: if train_config.precision is not None:
logger.info( logger.info(
"Setting precision float precision to {}", "Setting float32 matmul precision to {}",
train_config.precision, train_config.precision,
) )
torch.set_float32_matmul_precision(train_config.precision) torch.set_float32_matmul_precision(train_config.precision)
if train_config.compile_model:
logger.info("Compiling detector...")
compile_model(module.model)
logger.info("Starting main training loop...") logger.info("Starting main training loop...")
trainer.fit( trainer.fit(
module, module,

View File

@ -1,6 +1,6 @@
import uuid import uuid
from pathlib import Path from pathlib import Path
from typing import Callable, List, Optional from typing import Any, Callable, List, Optional, cast
from uuid import uuid4 from uuid import uuid4
import lightning as L import lightning as L
@ -15,6 +15,7 @@ from batdetect2.audio.clips import build_clipper
from batdetect2.audio.types import AudioLoader, ClipperProtocol from batdetect2.audio.types import AudioLoader, ClipperProtocol
from batdetect2.data import DatasetConfig, load_dataset from batdetect2.data import DatasetConfig, load_dataset
from batdetect2.data.annotations.batdetect2 import BatDetect2FilesAnnotations from batdetect2.data.annotations.batdetect2 import BatDetect2FilesAnnotations
from batdetect2.models.types import ModelProtocol
from batdetect2.preprocess import build_preprocessor from batdetect2.preprocess import build_preprocessor
from batdetect2.preprocess.types import PreprocessorProtocol from batdetect2.preprocess.types import PreprocessorProtocol
from batdetect2.targets import ( from batdetect2.targets import (
@ -156,12 +157,15 @@ def generate_whistle(tmp_path: Path):
offset = int((time - duration / 2) * samplerate) offset = int((time - duration / 2) * samplerate)
t = np.linspace(-duration / 2, duration / 2, frames, endpoint=False) t = np.linspace(-duration / 2, duration / 2, frames, endpoint=False)
data = signal.gausspulse( pulse = np.asarray(
t, signal.gausspulse(
fc=frequency, t,
bw=2 / (frequency * whistle_duration), fc=frequency,
bw=2 / (frequency * whistle_duration),
),
dtype=np.float64,
) )
wave = (np.roll(data, offset) * np.iinfo(np.int16).max).astype( wave = (np.roll(pulse, offset) * np.iinfo(np.int16).max).astype(
np.int16 np.int16
) )
sf.write(str(path), wave, samplerate, subtype="PCM_16") sf.write(str(path), wave, samplerate, subtype="PCM_16")
@ -363,6 +367,28 @@ def sample_audio_loader() -> AudioLoader:
return build_audio_loader() return build_audio_loader()
@pytest.fixture
def record_compiled_detector_calls(
monkeypatch: pytest.MonkeyPatch,
) -> Callable[[ModelProtocol], list[None]]:
def factory(model: ModelProtocol) -> list[None]:
compiled_calls: list[None] = []
detector = cast(Any, model.detector)
original_call_impl = detector._call_impl
def compile_detector() -> None:
def compiled_call(*args, **kwargs):
compiled_calls.append(None)
return original_call_impl(*args, **kwargs)
detector._compiled_call_impl = compiled_call
monkeypatch.setattr(detector, "compile", compile_detector)
return compiled_calls
return factory
@pytest.fixture @pytest.fixture
def bat_tag() -> data.Tag: def bat_tag() -> data.Tag:
return data.Tag(key="class", value="bat") return data.Tag(key="class", value="bat")

View File

@ -153,6 +153,40 @@ def test_process_spectrogram_rejects_batched_input(
api_v2.process_spectrogram(spec) api_v2.process_spectrogram(spec)
def test_user_can_compile_api_detector(
api_v2: BatDetect2API,
example_audio_files: list[Path],
record_compiled_detector_calls,
) -> None:
compiled_calls = record_compiled_detector_calls(api_v2.model)
audio = api_v2.load_audio(example_audio_files[0])
spec = api_v2.generate_spectrogram(audio)
api_v2.compile()
api_v2.compile()
api_v2.process_spectrogram(spec)
assert len(compiled_calls) == 1
def test_api_from_config_compiles_detector_when_requested(
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_config(
compile_model=True,
)
assert compiled_models == [api.model]
def test_user_can_read_top_class_and_other_class_scores( def test_user_can_read_top_class_and_other_class_scores(
api_v2: BatDetect2API, api_v2: BatDetect2API,
example_audio_files: list[Path], example_audio_files: list[Path],

View File

@ -3,6 +3,8 @@ from pathlib import Path
import pytest import pytest
from soundevent import data from soundevent import data
from batdetect2.api_v2 import BatDetect2API
from batdetect2.inference import InferenceConfig
from batdetect2.inference.batch import run_batch_inference from batdetect2.inference.batch import run_batch_inference
from batdetect2.targets import build_roi_mapping, build_targets from batdetect2.targets import build_roi_mapping, build_targets
from batdetect2.train import load_model_from_checkpoint from batdetect2.train import load_model_from_checkpoint
@ -53,3 +55,28 @@ def test_run_batch_inference_matches_single_clip_inference(
strict=True, strict=True,
): ):
assert_clip_detections_equal(batched, single) assert_clip_detections_equal(batched, single)
def test_run_batch_inference_compiles_detector_when_config_requests_compile(
example_annotations: list[data.ClipAnnotation],
record_compiled_detector_calls,
) -> None:
api = BatDetect2API.from_config()
compiled_calls = record_compiled_detector_calls(api.model)
api.compile()
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 compiled_calls

View File

@ -49,6 +49,16 @@ def build_default_module(
) )
def build_fast_train_config() -> TrainingConfig:
train_config = TrainingConfig()
train_config.trainer.limit_train_batches = 1
train_config.trainer.limit_val_batches = 1
train_config.trainer.log_every_n_steps = 1
train_config.train_loader.batch_size = 1
train_config.train_loader.augmentations.enabled = False
return train_config
def test_can_initialize_default_module(): def test_can_initialize_default_module():
module = build_default_module() module = build_default_module()
assert isinstance(module, L.LightningModule) assert isinstance(module, L.LightningModule)
@ -271,19 +281,7 @@ def test_train_smoke_produces_loadable_checkpoint(
sample_audio_loader: AudioLoader, sample_audio_loader: AudioLoader,
): ):
# Given # Given
train_config = TrainingConfig.model_validate( train_config = build_fast_train_config()
{
"trainer": {
"limit_train_batches": 1,
"limit_val_batches": 1,
"log_every_n_steps": 1,
},
"train_loader": {
"batch_size": 1,
"augmentations": {"enabled": False},
},
}
)
# When # When
run_train( run_train(
@ -310,6 +308,73 @@ def test_train_smoke_produces_loadable_checkpoint(
assert outputs is not None assert outputs is not None
@pytest.mark.slow
def test_run_train_compiles_detector_when_train_config_requests_compile(
tmp_path: Path,
example_annotations: list[data.ClipAnnotation],
record_compiled_detector_calls,
) -> None:
targets_config = TargetConfig()
targets = build_targets(targets_config)
roi_mapper = build_roi_mapping(targets_config.roi)
model = build_model(
ModelConfig(),
class_names=targets.class_names,
dimension_names=roi_mapper.dimension_names,
)
train_config = build_fast_train_config()
train_config.compile_model = True
compiled_calls = record_compiled_detector_calls(model)
module = run_train(
train_annotations=example_annotations[:1],
val_annotations=example_annotations[:1],
model=model,
targets=targets,
roi_mapper=roi_mapper,
targets_config=targets_config,
train_config=train_config,
num_epochs=1,
train_workers=0,
val_workers=0,
checkpoint_dir=tmp_path / "checkpoints",
log_dir=tmp_path / "logs",
seed=0,
)
assert (
getattr(module.model.detector, "_compiled_call_impl", None) is not None
)
assert compiled_calls
@pytest.mark.slow
def test_run_train_sets_float32_matmul_precision(
tmp_path: Path,
example_annotations: list[data.ClipAnnotation],
) -> None:
original_precision = torch.get_float32_matmul_precision()
train_config = build_fast_train_config()
train_config.precision = "high"
try:
run_train(
train_annotations=example_annotations[:1],
val_annotations=example_annotations[:1],
train_config=train_config,
num_epochs=1,
train_workers=0,
val_workers=0,
checkpoint_dir=tmp_path / "checkpoints",
log_dir=tmp_path / "logs",
seed=0,
)
assert torch.get_float32_matmul_precision() == "high"
finally:
torch.set_float32_matmul_precision(original_precision)
def test_build_training_module_uses_provided_model() -> None: def test_build_training_module_uses_provided_model() -> None:
targets = build_targets(TargetConfig()) targets = build_targets(TargetConfig())
roi_mapper = build_roi_mapping(TargetConfig().roi) roi_mapper = build_roi_mapping(TargetConfig().roi)