From 74922802185dbf19b290c202ac797ea8747f1bd1 Mon Sep 17 00:00:00 2001 From: Santiago Martinez Balvanera Date: Tue, 4 Aug 2026 08:45:29 +0100 Subject: [PATCH] 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,