Add setting for compiling the model before train

This commit is contained in:
Santiago Martinez Balvanera 2026-08-04 08:45:29 +01:00 committed by mbsantiago
parent d7896c6d75
commit 7492280218
3 changed files with 19 additions and 0 deletions

View File

@ -112,6 +112,9 @@ 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
@ -129,6 +132,7 @@ 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(

View File

@ -1,3 +1,5 @@
from typing import Literal
from pydantic import Field from pydantic import Field
from batdetect2.core.configs import BaseConfig from batdetect2.core.configs import BaseConfig
@ -39,6 +41,7 @@ class PLTrainerConfig(BaseConfig):
class TrainingConfig(BaseConfig): class TrainingConfig(BaseConfig):
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)
optimizer: OptimizerConfig = Field(default_factory=AdamOptimizerConfig) optimizer: OptimizerConfig = Field(default_factory=AdamOptimizerConfig)

View File

@ -2,6 +2,7 @@ from collections.abc import Sequence
from pathlib import Path from pathlib import Path
from typing import Optional from typing import Optional
import torch
from lightning import Trainer, seed_everything from lightning import Trainer, seed_everything
from lightning.pytorch.loggers import Logger from lightning.pytorch.loggers import Logger
from loguru import logger from loguru import logger
@ -206,6 +207,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:
logger.info(
"Setting precision float precision to {}",
train_config.precision,
)
torch.set_float32_matmul_precision(train_config.precision)
logger.info("Starting main training loop...") logger.info("Starting main training loop...")
trainer.fit( trainer.fit(
module, module,