From 13e6f97fa71492c6500dccda7db910c8b1daa63d Mon Sep 17 00:00:00 2001 From: Santiago Martinez Balvanera Date: Tue, 4 Aug 2026 08:45:29 +0100 Subject: [PATCH 1/3] Add setting for compiling the model before train --- src/batdetect2/models/__init__.py | 4 ++++ src/batdetect2/train/config.py | 3 +++ src/batdetect2/train/train.py | 12 ++++++++++++ 3 files changed, 19 insertions(+) diff --git a/src/batdetect2/models/__init__.py b/src/batdetect2/models/__init__.py index ee96d93..8564d05 100644 --- a/src/batdetect2/models/__init__.py +++ b/src/batdetect2/models/__init__.py @@ -112,6 +112,9 @@ class ModelConfig(BaseConfig): Attributes ---------- + compile : bool + If ``True``, compile the model before training. Defaults to + ``False``. samplerate : int Expected input audio sample rate in Hz. Audio must be resampled to this rate before being passed to the model. Defaults to @@ -129,6 +132,7 @@ class ModelConfig(BaseConfig): ``PostprocessConfig()``. """ + compile: bool = False samplerate: int = Field(default=TARGET_SAMPLERATE_HZ, gt=0) architecture: BackboneConfig = Field(default_factory=UNetBackboneConfig) preprocess: PreprocessingConfig = Field( diff --git a/src/batdetect2/train/config.py b/src/batdetect2/train/config.py index 94584cb..d12f332 100644 --- a/src/batdetect2/train/config.py +++ b/src/batdetect2/train/config.py @@ -1,3 +1,5 @@ +from typing import Literal + from pydantic import Field from batdetect2.core.configs import BaseConfig @@ -39,6 +41,7 @@ class PLTrainerConfig(BaseConfig): class TrainingConfig(BaseConfig): + precision: Literal["medium", "high"] | None = None train_loader: TrainLoaderConfig = Field(default_factory=TrainLoaderConfig) val_loader: ValLoaderConfig = Field(default_factory=ValLoaderConfig) optimizer: OptimizerConfig = Field(default_factory=AdamOptimizerConfig) diff --git a/src/batdetect2/train/train.py b/src/batdetect2/train/train.py index f59138b..ed3218b 100644 --- a/src/batdetect2/train/train.py +++ b/src/batdetect2/train/train.py @@ -2,6 +2,7 @@ from collections.abc import Sequence from pathlib import Path from typing import Optional +import torch from lightning import Trainer, seed_everything from lightning.pytorch.loggers import Logger from loguru import logger @@ -206,6 +207,17 @@ def run_train( run_name=run_name, ) + if model_config.compile: + logger.info("Compiling model...") + module.compile() + + if train_config.precision is not None: + logger.info( + "Setting precision float precision to {}", + train_config.precision, + ) + torch.set_float32_matmul_precision(train_config.precision) + logger.info("Starting main training loop...") trainer.fit( module, From 573a0ca39295771e62b5e55f21e04e65dcb10303 Mon Sep 17 00:00:00 2001 From: Santiago Martinez Balvanera Date: Tue, 4 Aug 2026 08:46:32 +0100 Subject: [PATCH 2/3] Allow passing the logger object directly to the train workflow --- src/batdetect2/api_v2.py | 3 +++ src/batdetect2/train/train.py | 3 ++- 2 files changed, 5 insertions(+), 1 deletion(-) diff --git a/src/batdetect2/api_v2.py b/src/batdetect2/api_v2.py index d2c0a8f..e84b876 100644 --- a/src/batdetect2/api_v2.py +++ b/src/batdetect2/api_v2.py @@ -1,4 +1,5 @@ from __future__ import annotations +from lightning.pytorch.loggers import Logger from pathlib import Path from typing import TYPE_CHECKING, Literal @@ -192,6 +193,7 @@ class BatDetect2API: train_config: TrainingConfig | None = None, logger_config: LoggerConfig | None = None, logging_callbacks: Sequence[LoggingCallback[TrainLoggingContext]] = (), + train_logger: Logger | None = None, ): """Train the current model on a set of annotations. @@ -255,6 +257,7 @@ class BatDetect2API: audio_config=audio_config or self.audio_config, logger_config=logger_config or self.logging_config.train, logging_callbacks=logging_callbacks, + train_logger=train_logger, ) self.model.eval() return self diff --git a/src/batdetect2/train/train.py b/src/batdetect2/train/train.py index ed3218b..0f4d8aa 100644 --- a/src/batdetect2/train/train.py +++ b/src/batdetect2/train/train.py @@ -72,6 +72,7 @@ def run_train( run_name: str | None = None, seed: int | None = None, logging_callbacks: Sequence[LoggingCallback[TrainLoggingContext]] = (), + train_logger: Logger | None = None, ): if seed is not None: seed_everything(seed) @@ -165,7 +166,7 @@ def run_train( roi_mapper=roi_mapper, ) - train_logger = build_logger( + train_logger = train_logger or build_logger( logger_config or CSVLoggerConfig(), log_dir=log_dir, experiment_name=experiment_name, From 1bab80f8c6cd880adfc306d4632e582ddfe45653 Mon Sep 17 00:00:00 2001 From: Santiago Martinez Balvanera Date: Tue, 4 Aug 2026 08:46:51 +0100 Subject: [PATCH 3/3] Expand the cosine annealing config --- src/batdetect2/train/schedulers.py | 39 ++++++++++++++++++------------ 1 file changed, 23 insertions(+), 16 deletions(-) diff --git a/src/batdetect2/train/schedulers.py b/src/batdetect2/train/schedulers.py index 73ebd09..a21ef5b 100644 --- a/src/batdetect2/train/schedulers.py +++ b/src/batdetect2/train/schedulers.py @@ -23,21 +23,6 @@ __all__ = [ ] -class CosineAnnealingSchedulerConfig(BaseConfig): - """Configuration for ``CosineAnnealingLR``. - - Attributes - ---------- - name : Literal["cosine_annealing"] - Discriminator field used by the scheduler registry. - t_max : int - Number of epochs to complete one cosine cycle. - """ - - name: Literal["cosine_annealing"] = "cosine_annealing" - t_max: int = 200 - - scheduler_registry: Registry[LRScheduler, [Optimizer]] = Registry("scheduler") @@ -53,6 +38,24 @@ class SchedulerImportConfig(ImportConfig): name: Literal["import"] = "import" +class CosineAnnealingSchedulerConfig(BaseConfig): + """Configuration for ``CosineAnnealingLR``. + + Attributes + ---------- + name : Literal["cosine_annealing"] + Discriminator field used by the scheduler registry. + t_max : int + Number of epochs to complete one cosine cycle. + eta_min : float, optional + Minimum learning rate. Defaults to 0. + """ + + name: Literal["cosine_annealing"] = "cosine_annealing" + t_max: int = 200 + eta_min: float = 0 + + @scheduler_registry.register(CosineAnnealingSchedulerConfig) def build_cosine_scheduler( config: CosineAnnealingSchedulerConfig, @@ -63,7 +66,11 @@ def build_cosine_scheduler( ``t_max`` is interpreted in epochs because Lightning steps the scheduler once per epoch when ``interval="epoch"`` is used. """ - return CosineAnnealingLR(optimizer, T_max=config.t_max) + return CosineAnnealingLR( + optimizer, + T_max=config.t_max, + eta_min=config.eta_min, + ) SchedulerConfig = Annotated[