mirror of
https://github.com/macaodha/batdetect2.git
synced 2026-08-22 03:00:10 +02:00
Allow passing the logger object directly to the train workflow
This commit is contained in:
parent
13e6f97fa7
commit
573a0ca392
@ -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
|
||||
|
||||
@ -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,
|
||||
|
||||
Loading…
Reference in New Issue
Block a user