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
7492280218
commit
ead0adf284
@ -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
|
||||||
|
|||||||
@ -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,
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user