mirror of
https://github.com/macaodha/batdetect2.git
synced 2026-08-22 03:00:10 +02:00
Add setting for compiling the model before train
This commit is contained in:
parent
d7896c6d75
commit
7492280218
@ -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(
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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,
|
||||
|
||||
Loading…
Reference in New Issue
Block a user