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 committed by mbsantiago
parent 7492280218
commit ead0adf284
2 changed files with 5 additions and 1 deletions

View File

@ -1,4 +1,5 @@
from __future__ import annotations from __future__ import annotations
from lightning.pytorch.loggers import Logger
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Literal from typing import TYPE_CHECKING, Literal
@ -192,6 +193,7 @@ class BatDetect2API:
train_config: TrainingConfig | None = None, train_config: TrainingConfig | None = None,
logger_config: LoggerConfig | None = None, logger_config: LoggerConfig | None = None,
logging_callbacks: Sequence[LoggingCallback[TrainLoggingContext]] = (), logging_callbacks: Sequence[LoggingCallback[TrainLoggingContext]] = (),
train_logger: Logger | None = None,
): ):
"""Train the current model on a set of annotations. """Train the current model on a set of annotations.
@ -255,6 +257,7 @@ class BatDetect2API:
audio_config=audio_config or self.audio_config, audio_config=audio_config or self.audio_config,
logger_config=logger_config or self.logging_config.train, logger_config=logger_config or self.logging_config.train,
logging_callbacks=logging_callbacks, logging_callbacks=logging_callbacks,
train_logger=train_logger,
) )
self.model.eval() self.model.eval()
return self return self

View File

@ -72,6 +72,7 @@ def run_train(
run_name: str | None = None, run_name: str | None = None,
seed: int | None = None, seed: int | None = None,
logging_callbacks: Sequence[LoggingCallback[TrainLoggingContext]] = (), logging_callbacks: Sequence[LoggingCallback[TrainLoggingContext]] = (),
train_logger: Logger | None = None,
): ):
if seed is not None: if seed is not None:
seed_everything(seed) seed_everything(seed)
@ -165,7 +166,7 @@ def run_train(
roi_mapper=roi_mapper, roi_mapper=roi_mapper,
) )
train_logger = build_logger( train_logger = train_logger or build_logger(
logger_config or CSVLoggerConfig(), logger_config or CSVLoggerConfig(),
log_dir=log_dir, log_dir=log_dir,
experiment_name=experiment_name, experiment_name=experiment_name,