Merge pull request #12763 from freqtrade/feat/tb-lightgbm
feat: Add tensorboard callback to lightgbm
This commit is contained in:
@@ -1,10 +1,12 @@
|
|||||||
import logging
|
import logging
|
||||||
|
from collections.abc import Callable
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from lightgbm import LGBMClassifier
|
from lightgbm import LGBMClassifier
|
||||||
|
|
||||||
from freqtrade.freqai.base_models.BaseClassifierModel import BaseClassifierModel
|
from freqtrade.freqai.base_models.BaseClassifierModel import BaseClassifierModel
|
||||||
from freqtrade.freqai.data_kitchen import FreqaiDataKitchen
|
from freqtrade.freqai.data_kitchen import FreqaiDataKitchen
|
||||||
|
from freqtrade.freqai.tensorboard import LightGBMCallback
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -46,6 +48,10 @@ class LightGBMClassifier(BaseClassifierModel):
|
|||||||
init_model = self.get_init_model(dk.pair)
|
init_model = self.get_init_model(dk.pair)
|
||||||
|
|
||||||
model = LGBMClassifier(**self.model_training_parameters)
|
model = LGBMClassifier(**self.model_training_parameters)
|
||||||
|
activate_tensorboard = self.freqai_info.get("activate_tensorboard", True)
|
||||||
|
callbacks: list[Callable[..., Any]] = []
|
||||||
|
if LightGBMCallback is not None:
|
||||||
|
callbacks = [LightGBMCallback(dk.data_path, activate_tensorboard)]
|
||||||
model.fit(
|
model.fit(
|
||||||
X=X,
|
X=X,
|
||||||
y=y,
|
y=y,
|
||||||
@@ -53,6 +59,7 @@ class LightGBMClassifier(BaseClassifierModel):
|
|||||||
sample_weight=train_weights,
|
sample_weight=train_weights,
|
||||||
eval_sample_weight=[test_weights],
|
eval_sample_weight=[test_weights],
|
||||||
init_model=init_model,
|
init_model=init_model,
|
||||||
|
callbacks=callbacks,
|
||||||
)
|
)
|
||||||
|
|
||||||
return model
|
return model
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ from lightgbm import LGBMClassifier
|
|||||||
from freqtrade.freqai.base_models.BaseClassifierModel import BaseClassifierModel
|
from freqtrade.freqai.base_models.BaseClassifierModel import BaseClassifierModel
|
||||||
from freqtrade.freqai.base_models.FreqaiMultiOutputClassifier import FreqaiMultiOutputClassifier
|
from freqtrade.freqai.base_models.FreqaiMultiOutputClassifier import FreqaiMultiOutputClassifier
|
||||||
from freqtrade.freqai.data_kitchen import FreqaiDataKitchen
|
from freqtrade.freqai.data_kitchen import FreqaiDataKitchen
|
||||||
|
from freqtrade.freqai.tensorboard import LightGBMCallback
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -53,6 +54,11 @@ class LightGBMClassifierMultiTarget(BaseClassifierModel):
|
|||||||
else:
|
else:
|
||||||
init_models = [None] * y.shape[1]
|
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 = []
|
fit_params = []
|
||||||
for i in range(len(eval_sets)):
|
for i in range(len(eval_sets)):
|
||||||
fit_params.append(
|
fit_params.append(
|
||||||
@@ -60,6 +66,7 @@ class LightGBMClassifierMultiTarget(BaseClassifierModel):
|
|||||||
"eval_set": eval_sets[i],
|
"eval_set": eval_sets[i],
|
||||||
"eval_sample_weight": eval_weights,
|
"eval_sample_weight": eval_weights,
|
||||||
"init_model": init_models[i],
|
"init_model": init_models[i],
|
||||||
|
"callbacks": callbacks,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1,10 +1,12 @@
|
|||||||
import logging
|
import logging
|
||||||
|
from collections.abc import Callable
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from lightgbm import LGBMRegressor
|
from lightgbm import LGBMRegressor
|
||||||
|
|
||||||
from freqtrade.freqai.base_models.BaseRegressionModel import BaseRegressionModel
|
from freqtrade.freqai.base_models.BaseRegressionModel import BaseRegressionModel
|
||||||
from freqtrade.freqai.data_kitchen import FreqaiDataKitchen
|
from freqtrade.freqai.data_kitchen import FreqaiDataKitchen
|
||||||
|
from freqtrade.freqai.tensorboard import LightGBMCallback
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -42,6 +44,11 @@ class LightGBMRegressor(BaseRegressionModel):
|
|||||||
|
|
||||||
model = LGBMRegressor(**self.model_training_parameters)
|
model = LGBMRegressor(**self.model_training_parameters)
|
||||||
|
|
||||||
|
activate_tensorboard = self.freqai_info.get("activate_tensorboard", True)
|
||||||
|
callbacks: list[Callable[..., Any]] = []
|
||||||
|
if LightGBMCallback is not None:
|
||||||
|
callbacks = [LightGBMCallback(dk.data_path, activate_tensorboard)]
|
||||||
|
|
||||||
model.fit(
|
model.fit(
|
||||||
X=X,
|
X=X,
|
||||||
y=y,
|
y=y,
|
||||||
@@ -49,6 +56,7 @@ class LightGBMRegressor(BaseRegressionModel):
|
|||||||
sample_weight=train_weights,
|
sample_weight=train_weights,
|
||||||
eval_sample_weight=[eval_weights],
|
eval_sample_weight=[eval_weights],
|
||||||
init_model=init_model,
|
init_model=init_model,
|
||||||
|
callbacks=callbacks,
|
||||||
)
|
)
|
||||||
|
|
||||||
return model
|
return model
|
||||||
|
|||||||
@@ -6,6 +6,7 @@ from lightgbm import LGBMRegressor
|
|||||||
from freqtrade.freqai.base_models.BaseRegressionModel import BaseRegressionModel
|
from freqtrade.freqai.base_models.BaseRegressionModel import BaseRegressionModel
|
||||||
from freqtrade.freqai.base_models.FreqaiMultiOutputRegressor import FreqaiMultiOutputRegressor
|
from freqtrade.freqai.base_models.FreqaiMultiOutputRegressor import FreqaiMultiOutputRegressor
|
||||||
from freqtrade.freqai.data_kitchen import FreqaiDataKitchen
|
from freqtrade.freqai.data_kitchen import FreqaiDataKitchen
|
||||||
|
from freqtrade.freqai.tensorboard import LightGBMCallback
|
||||||
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -55,6 +56,11 @@ class LightGBMRegressorMultiTarget(BaseRegressionModel):
|
|||||||
else:
|
else:
|
||||||
init_models = [None] * y.shape[1]
|
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 = []
|
fit_params = []
|
||||||
for i in range(len(eval_sets)):
|
for i in range(len(eval_sets)):
|
||||||
fit_params.append(
|
fit_params.append(
|
||||||
@@ -62,6 +68,7 @@ class LightGBMRegressorMultiTarget(BaseRegressionModel):
|
|||||||
"eval_set": eval_sets[i],
|
"eval_set": eval_sets[i],
|
||||||
"eval_sample_weight": eval_weights,
|
"eval_sample_weight": eval_weights,
|
||||||
"init_model": init_models[i],
|
"init_model": init_models[i],
|
||||||
|
"callbacks": callbacks,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -1,9 +1,11 @@
|
|||||||
# ensure users can still use a non-torch freqai version
|
# ensure users can still use a non-torch freqai version
|
||||||
try:
|
try:
|
||||||
|
from freqtrade.freqai.tensorboard.lightgbm_callback import LightGBMTensorboardCallback
|
||||||
from freqtrade.freqai.tensorboard.tensorboard import TensorBoardCallback, TensorboardLogger
|
from freqtrade.freqai.tensorboard.tensorboard import TensorBoardCallback, TensorboardLogger
|
||||||
|
|
||||||
TBLogger = TensorboardLogger
|
TBLogger = TensorboardLogger
|
||||||
TBCallback = TensorBoardCallback
|
TBCallback = TensorBoardCallback
|
||||||
|
LightGBMCallback = LightGBMTensorboardCallback
|
||||||
except ModuleNotFoundError:
|
except ModuleNotFoundError:
|
||||||
from freqtrade.freqai.tensorboard.base_tensorboard import (
|
from freqtrade.freqai.tensorboard.base_tensorboard import (
|
||||||
BaseTensorBoardCallback,
|
BaseTensorBoardCallback,
|
||||||
@@ -12,5 +14,6 @@ except ModuleNotFoundError:
|
|||||||
|
|
||||||
TBLogger = BaseTensorboardLogger # type: ignore
|
TBLogger = BaseTensorboardLogger # type: ignore
|
||||||
TBCallback = BaseTensorBoardCallback # type: ignore
|
TBCallback = BaseTensorBoardCallback # type: ignore
|
||||||
|
LightGBMCallback = None # type: ignore
|
||||||
|
|
||||||
__all__ = ("TBLogger", "TBCallback")
|
__all__ = ("TBLogger", "TBCallback", "LightGBMCallback")
|
||||||
|
|||||||
@@ -0,0 +1,24 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from freqtrade.freqai.tensorboard.tensorboard import TensorboardLogger
|
||||||
|
|
||||||
|
|
||||||
|
class LightGBMTensorboardCallback:
|
||||||
|
def __init__(self, logdir, activate: bool) -> None:
|
||||||
|
self.activate = activate
|
||||||
|
self.logger = TensorboardLogger(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()
|
||||||
Reference in New Issue
Block a user