diff --git a/src/batdetect2/api_v2.py b/src/batdetect2/api_v2.py index d2c0a8f..e84b876 100644 --- a/src/batdetect2/api_v2.py +++ b/src/batdetect2/api_v2.py @@ -1,4 +1,5 @@ from __future__ import annotations +from lightning.pytorch.loggers import Logger from pathlib import Path from typing import TYPE_CHECKING, Literal @@ -192,6 +193,7 @@ class BatDetect2API: train_config: TrainingConfig | None = None, logger_config: LoggerConfig | None = None, logging_callbacks: Sequence[LoggingCallback[TrainLoggingContext]] = (), + train_logger: Logger | None = None, ): """Train the current model on a set of annotations. @@ -255,6 +257,7 @@ class BatDetect2API: audio_config=audio_config or self.audio_config, logger_config=logger_config or self.logging_config.train, logging_callbacks=logging_callbacks, + train_logger=train_logger, ) self.model.eval() return self diff --git a/src/batdetect2/train/train.py b/src/batdetect2/train/train.py index ed3218b..0f4d8aa 100644 --- a/src/batdetect2/train/train.py +++ b/src/batdetect2/train/train.py @@ -72,6 +72,7 @@ def run_train( run_name: str | None = None, seed: int | None = None, logging_callbacks: Sequence[LoggingCallback[TrainLoggingContext]] = (), + train_logger: Logger | None = None, ): if seed is not None: seed_everything(seed) @@ -165,7 +166,7 @@ def run_train( roi_mapper=roi_mapper, ) - train_logger = build_logger( + train_logger = train_logger or build_logger( logger_config or CSVLoggerConfig(), log_dir=log_dir, experiment_name=experiment_name,