From 455954e0e3d79a82c58bd9ae4801a811c86d7a61 Mon Sep 17 00:00:00 2001 From: Robert Caulk Date: Sun, 1 Feb 2026 15:14:48 +0100 Subject: [PATCH 1/3] feat: Add tensorboard callback to lightgbm --- .../prediction_models/LightGBMClassifier.py | 6 +++++ .../LightGBMClassifierMultiTarget.py | 7 ++++++ .../prediction_models/LightGBMRegressor.py | 7 ++++++ .../LightGBMRegressorMultiTarget.py | 7 ++++++ freqtrade/freqai/tensorboard/__init__.py | 5 +++- .../freqai/tensorboard/lightgbm_callback.py | 24 +++++++++++++++++++ 6 files changed, 55 insertions(+), 1 deletion(-) create mode 100644 freqtrade/freqai/tensorboard/lightgbm_callback.py diff --git a/freqtrade/freqai/prediction_models/LightGBMClassifier.py b/freqtrade/freqai/prediction_models/LightGBMClassifier.py index e17f7417c..75bbe2d1a 100644 --- a/freqtrade/freqai/prediction_models/LightGBMClassifier.py +++ b/freqtrade/freqai/prediction_models/LightGBMClassifier.py @@ -5,6 +5,7 @@ from lightgbm import LGBMClassifier from freqtrade.freqai.base_models.BaseClassifierModel import BaseClassifierModel from freqtrade.freqai.data_kitchen import FreqaiDataKitchen +from freqtrade.freqai.tensorboard import LightGBMCallback logger = logging.getLogger(__name__) @@ -46,6 +47,10 @@ class LightGBMClassifier(BaseClassifierModel): init_model = self.get_init_model(dk.pair) model = LGBMClassifier(**self.model_training_parameters) + activate_tensorboard = self.freqai_info.get("activate_tensorboard", True) + callbacks = [] + if LightGBMCallback is not None: + callbacks = [LightGBMCallback(dk.data_path, activate_tensorboard)] model.fit( X=X, y=y, @@ -53,6 +58,7 @@ class LightGBMClassifier(BaseClassifierModel): sample_weight=train_weights, eval_sample_weight=[test_weights], init_model=init_model, + callbacks=callbacks, ) return model diff --git a/freqtrade/freqai/prediction_models/LightGBMClassifierMultiTarget.py b/freqtrade/freqai/prediction_models/LightGBMClassifierMultiTarget.py index 9fb775614..4e4d981b8 100644 --- a/freqtrade/freqai/prediction_models/LightGBMClassifierMultiTarget.py +++ b/freqtrade/freqai/prediction_models/LightGBMClassifierMultiTarget.py @@ -6,6 +6,7 @@ from lightgbm import LGBMClassifier from freqtrade.freqai.base_models.BaseClassifierModel import BaseClassifierModel from freqtrade.freqai.base_models.FreqaiMultiOutputClassifier import FreqaiMultiOutputClassifier from freqtrade.freqai.data_kitchen import FreqaiDataKitchen +from freqtrade.freqai.tensorboard import LightGBMCallback logger = logging.getLogger(__name__) @@ -53,6 +54,11 @@ class LightGBMClassifierMultiTarget(BaseClassifierModel): else: init_models = [None] * y.shape[1] + activate_tensorboard = self.freqai_info.get("activate_tensorboard", True) + callbacks = [] + if LightGBMCallback is not None: + callbacks = [LightGBMCallback(dk.data_path, activate_tensorboard)] + fit_params = [] for i in range(len(eval_sets)): fit_params.append( @@ -60,6 +66,7 @@ class LightGBMClassifierMultiTarget(BaseClassifierModel): "eval_set": eval_sets[i], "eval_sample_weight": eval_weights, "init_model": init_models[i], + "callbacks": callbacks, } ) diff --git a/freqtrade/freqai/prediction_models/LightGBMRegressor.py b/freqtrade/freqai/prediction_models/LightGBMRegressor.py index d55cd0ca2..af89d4825 100644 --- a/freqtrade/freqai/prediction_models/LightGBMRegressor.py +++ b/freqtrade/freqai/prediction_models/LightGBMRegressor.py @@ -5,6 +5,7 @@ from lightgbm import LGBMRegressor from freqtrade.freqai.base_models.BaseRegressionModel import BaseRegressionModel from freqtrade.freqai.data_kitchen import FreqaiDataKitchen +from freqtrade.freqai.tensorboard import LightGBMCallback logger = logging.getLogger(__name__) @@ -42,6 +43,11 @@ class LightGBMRegressor(BaseRegressionModel): model = LGBMRegressor(**self.model_training_parameters) + activate_tensorboard = self.freqai_info.get("activate_tensorboard", True) + callbacks = [] + if LightGBMCallback is not None: + callbacks = [LightGBMCallback(dk.data_path, activate_tensorboard)] + model.fit( X=X, y=y, @@ -49,6 +55,7 @@ class LightGBMRegressor(BaseRegressionModel): sample_weight=train_weights, eval_sample_weight=[eval_weights], init_model=init_model, + callbacks=callbacks, ) return model diff --git a/freqtrade/freqai/prediction_models/LightGBMRegressorMultiTarget.py b/freqtrade/freqai/prediction_models/LightGBMRegressorMultiTarget.py index c4669a79d..8f374190b 100644 --- a/freqtrade/freqai/prediction_models/LightGBMRegressorMultiTarget.py +++ b/freqtrade/freqai/prediction_models/LightGBMRegressorMultiTarget.py @@ -6,6 +6,7 @@ from lightgbm import LGBMRegressor from freqtrade.freqai.base_models.BaseRegressionModel import BaseRegressionModel from freqtrade.freqai.base_models.FreqaiMultiOutputRegressor import FreqaiMultiOutputRegressor from freqtrade.freqai.data_kitchen import FreqaiDataKitchen +from freqtrade.freqai.tensorboard import LightGBMCallback logger = logging.getLogger(__name__) @@ -55,6 +56,11 @@ class LightGBMRegressorMultiTarget(BaseRegressionModel): else: init_models = [None] * y.shape[1] + activate_tensorboard = self.freqai_info.get("activate_tensorboard", True) + callbacks = [] + if LightGBMCallback is not None: + callbacks = [LightGBMCallback(dk.data_path, activate_tensorboard)] + fit_params = [] for i in range(len(eval_sets)): fit_params.append( @@ -62,6 +68,7 @@ class LightGBMRegressorMultiTarget(BaseRegressionModel): "eval_set": eval_sets[i], "eval_sample_weight": eval_weights, "init_model": init_models[i], + "callbacks": callbacks, } ) diff --git a/freqtrade/freqai/tensorboard/__init__.py b/freqtrade/freqai/tensorboard/__init__.py index 183c25b22..68d045bf8 100644 --- a/freqtrade/freqai/tensorboard/__init__.py +++ b/freqtrade/freqai/tensorboard/__init__.py @@ -1,9 +1,11 @@ # ensure users can still use a non-torch freqai version try: + from freqtrade.freqai.tensorboard.lightgbm_callback import LightGBMTensorboardCallback from freqtrade.freqai.tensorboard.tensorboard import TensorBoardCallback, TensorboardLogger TBLogger = TensorboardLogger TBCallback = TensorBoardCallback + LightGBMCallback = LightGBMTensorboardCallback except ModuleNotFoundError: from freqtrade.freqai.tensorboard.base_tensorboard import ( BaseTensorBoardCallback, @@ -12,5 +14,6 @@ except ModuleNotFoundError: TBLogger = BaseTensorboardLogger # type: ignore TBCallback = BaseTensorBoardCallback # type: ignore + LightGBMCallback = None # type: ignore -__all__ = ("TBLogger", "TBCallback") +__all__ = ("TBLogger", "TBCallback", "LightGBMCallback") diff --git a/freqtrade/freqai/tensorboard/lightgbm_callback.py b/freqtrade/freqai/tensorboard/lightgbm_callback.py new file mode 100644 index 000000000..4e9bec804 --- /dev/null +++ b/freqtrade/freqai/tensorboard/lightgbm_callback.py @@ -0,0 +1,24 @@ +from __future__ import annotations + +from freqtrade.freqai.tensorboard import TBLogger + + +class LightGBMTensorboardCallback: + def __init__(self, logdir, activate: bool) -> None: + self.activate = activate + self.logger = TBLogger(logdir, activate) + + def __call__(self, env) -> None: + if not self.activate: + return + + evals = getattr(env, "evaluation_result_list", None) + if not evals: + return + + for data_name, metric_name, value, _ in evals: + self.logger.log_scalar(f"{data_name}-{metric_name}", value, env.iteration) + + end_iteration = getattr(env, "end_iteration", None) + if end_iteration is not None and env.iteration + 1 >= end_iteration: + self.logger.close() From 783c365c10c1976cd22ea44747e15c45369df6aa Mon Sep 17 00:00:00 2001 From: Robert Caulk Date: Thu, 5 Feb 2026 17:15:00 +0100 Subject: [PATCH 2/3] fix: Avoid circular import --- freqtrade/freqai/tensorboard/lightgbm_callback.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/freqtrade/freqai/tensorboard/lightgbm_callback.py b/freqtrade/freqai/tensorboard/lightgbm_callback.py index 4e9bec804..c71a43ec5 100644 --- a/freqtrade/freqai/tensorboard/lightgbm_callback.py +++ b/freqtrade/freqai/tensorboard/lightgbm_callback.py @@ -1,12 +1,12 @@ from __future__ import annotations -from freqtrade.freqai.tensorboard import TBLogger +from freqtrade.freqai.tensorboard.tensorboard import TensorboardLogger class LightGBMTensorboardCallback: def __init__(self, logdir, activate: bool) -> None: self.activate = activate - self.logger = TBLogger(logdir, activate) + self.logger = TensorboardLogger(logdir, activate) def __call__(self, env) -> None: if not self.activate: From 6814bc84aadb79419e37f6c3e83a5ad626c556ae Mon Sep 17 00:00:00 2001 From: Robert Caulk Date: Thu, 5 Feb 2026 17:24:17 +0100 Subject: [PATCH 3/3] fix: Try to fix linting --- freqtrade/freqai/prediction_models/LightGBMClassifier.py | 3 ++- freqtrade/freqai/prediction_models/LightGBMRegressor.py | 3 ++- 2 files changed, 4 insertions(+), 2 deletions(-) diff --git a/freqtrade/freqai/prediction_models/LightGBMClassifier.py b/freqtrade/freqai/prediction_models/LightGBMClassifier.py index 75bbe2d1a..0b78c4129 100644 --- a/freqtrade/freqai/prediction_models/LightGBMClassifier.py +++ b/freqtrade/freqai/prediction_models/LightGBMClassifier.py @@ -1,4 +1,5 @@ import logging +from collections.abc import Callable from typing import Any from lightgbm import LGBMClassifier @@ -48,7 +49,7 @@ class LightGBMClassifier(BaseClassifierModel): model = LGBMClassifier(**self.model_training_parameters) activate_tensorboard = self.freqai_info.get("activate_tensorboard", True) - callbacks = [] + callbacks: list[Callable[..., Any]] = [] if LightGBMCallback is not None: callbacks = [LightGBMCallback(dk.data_path, activate_tensorboard)] model.fit( diff --git a/freqtrade/freqai/prediction_models/LightGBMRegressor.py b/freqtrade/freqai/prediction_models/LightGBMRegressor.py index af89d4825..abd838eee 100644 --- a/freqtrade/freqai/prediction_models/LightGBMRegressor.py +++ b/freqtrade/freqai/prediction_models/LightGBMRegressor.py @@ -1,4 +1,5 @@ import logging +from collections.abc import Callable from typing import Any from lightgbm import LGBMRegressor @@ -44,7 +45,7 @@ class LightGBMRegressor(BaseRegressionModel): model = LGBMRegressor(**self.model_training_parameters) activate_tensorboard = self.freqai_info.get("activate_tensorboard", True) - callbacks = [] + callbacks: list[Callable[..., Any]] = [] if LightGBMCallback is not None: callbacks = [LightGBMCallback(dk.data_path, activate_tensorboard)]