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
0431ef2679
commit
13e6f97fa7
@ -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(
|
||||||
|
|||||||
@ -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)
|
||||||
|
|||||||
@ -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,
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user