Allow passing the logger object directly to the train workflow

This commit is contained in:
Santiago Martinez Balvanera 2026-08-04 08:46:32 +01:00
parent 13e6f97fa7
commit 573a0ca392
2 changed files with 5 additions and 1 deletions

View File

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

View File

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