Add setting for compiling the model before train

This commit is contained in:
Santiago Martinez Balvanera 2026-08-04 08:45:29 +01:00
parent 0431ef2679
commit 13e6f97fa7
3 changed files with 19 additions and 0 deletions

View File

@ -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(

View File

@ -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)

View File

@ -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,