2022-08-20 14:35:29 +00:00
|
|
|
import logging
|
2022-08-28 17:21:57 +00:00
|
|
|
from pathlib import Path
|
2022-08-20 14:35:29 +00:00
|
|
|
from typing import Any, Dict # , Tuple
|
|
|
|
|
|
|
|
# import numpy.typing as npt
|
|
|
|
import torch as th
|
2022-09-23 17:30:56 +00:00
|
|
|
from pandas import DataFrame
|
2022-08-20 14:35:29 +00:00
|
|
|
from stable_baselines3.common.callbacks import EvalCallback
|
|
|
|
from stable_baselines3.common.vec_env import SubprocVecEnv
|
2022-09-23 17:30:56 +00:00
|
|
|
|
2022-08-28 17:21:57 +00:00
|
|
|
from freqtrade.freqai.data_kitchen import FreqaiDataKitchen
|
2022-08-20 14:35:29 +00:00
|
|
|
from freqtrade.freqai.RL.BaseReinforcementLearningModel import (BaseReinforcementLearningModel,
|
|
|
|
make_env)
|
|
|
|
|
|
|
|
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
|
|
|
|
|
|
class ReinforcementLearner_multiproc(BaseReinforcementLearningModel):
|
|
|
|
"""
|
|
|
|
User created Reinforcement Learning Model prediction model.
|
|
|
|
"""
|
|
|
|
|
2022-09-14 22:46:35 +00:00
|
|
|
def fit(self, data_dictionary: Dict[str, Any], dk: FreqaiDataKitchen, **kwargs):
|
2022-08-20 14:35:29 +00:00
|
|
|
|
|
|
|
train_df = data_dictionary["train_features"]
|
|
|
|
total_timesteps = self.freqai_info["rl_config"]["train_cycles"] * len(train_df)
|
|
|
|
|
|
|
|
# model arch
|
|
|
|
policy_kwargs = dict(activation_fn=th.nn.ReLU,
|
2022-10-08 10:10:38 +00:00
|
|
|
net_arch=self.net_arch)
|
2022-08-20 14:35:29 +00:00
|
|
|
|
2022-08-25 09:46:18 +00:00
|
|
|
if dk.pair not in self.dd.model_dictionary or not self.continual_learning:
|
|
|
|
model = self.MODELCLASS(self.policy_type, self.train_env, policy_kwargs=policy_kwargs,
|
2022-08-31 14:50:39 +00:00
|
|
|
tensorboard_log=Path(
|
|
|
|
dk.full_path / "tensorboard" / dk.pair.split('/')[0]),
|
2022-08-25 09:46:18 +00:00
|
|
|
**self.freqai_info['model_training_parameters']
|
|
|
|
)
|
|
|
|
else:
|
2022-08-25 17:05:51 +00:00
|
|
|
logger.info('Continual learning activated - starting training from previously '
|
2022-08-25 09:46:18 +00:00
|
|
|
'trained agent.')
|
|
|
|
model = self.dd.model_dictionary[dk.pair]
|
|
|
|
model.set_env(self.train_env)
|
2022-08-20 14:35:29 +00:00
|
|
|
|
|
|
|
model.learn(
|
|
|
|
total_timesteps=int(total_timesteps),
|
|
|
|
callback=self.eval_callback
|
|
|
|
)
|
|
|
|
|
|
|
|
if Path(dk.data_path / "best_model.zip").is_file():
|
|
|
|
logger.info('Callback found a best model.')
|
|
|
|
best_model = self.MODELCLASS.load(dk.data_path / "best_model")
|
|
|
|
return best_model
|
|
|
|
|
|
|
|
logger.info('Couldnt find best model, using final model instead.')
|
|
|
|
|
|
|
|
return model
|
|
|
|
|
2022-09-23 17:17:27 +00:00
|
|
|
def set_train_and_eval_environments(self, data_dictionary: Dict[str, Any],
|
|
|
|
prices_train: DataFrame, prices_test: DataFrame,
|
|
|
|
dk: FreqaiDataKitchen):
|
2022-08-20 14:35:29 +00:00
|
|
|
"""
|
2022-09-23 17:17:27 +00:00
|
|
|
User can override this if they are using a custom MyRLEnv
|
2022-11-13 16:43:52 +00:00
|
|
|
:param data_dictionary: dict = common data dictionary containing train and test
|
2022-09-23 17:17:27 +00:00
|
|
|
features/labels/weights.
|
2022-11-13 16:43:52 +00:00
|
|
|
:param prices_train/test: DataFrame = dataframe comprised of the prices to be used in
|
2022-09-23 17:17:27 +00:00
|
|
|
the environment during training
|
|
|
|
or testing
|
2022-11-13 16:43:52 +00:00
|
|
|
:param dk: FreqaiDataKitchen = the datakitchen for the current pair
|
2022-08-20 14:35:29 +00:00
|
|
|
"""
|
|
|
|
train_df = data_dictionary["train_features"]
|
|
|
|
test_df = data_dictionary["test_features"]
|
|
|
|
|
2022-08-25 09:46:18 +00:00
|
|
|
env_id = "train_env"
|
2022-08-25 17:05:51 +00:00
|
|
|
self.train_env = SubprocVecEnv([make_env(self.MyRLEnv, env_id, i, 1, train_df, prices_train,
|
2022-08-28 17:21:57 +00:00
|
|
|
self.reward_params, self.CONV_WIDTH, monitor=True,
|
2022-08-25 09:46:18 +00:00
|
|
|
config=self.config) for i
|
2022-09-28 22:10:18 +00:00
|
|
|
in range(self.max_threads)])
|
2022-08-25 09:46:18 +00:00
|
|
|
|
|
|
|
eval_env_id = 'eval_env'
|
2022-08-25 17:05:51 +00:00
|
|
|
self.eval_env = SubprocVecEnv([make_env(self.MyRLEnv, eval_env_id, i, 1,
|
|
|
|
test_df, prices_test,
|
2022-08-25 09:46:18 +00:00
|
|
|
self.reward_params, self.CONV_WIDTH, monitor=True,
|
|
|
|
config=self.config) for i
|
2022-09-28 22:10:18 +00:00
|
|
|
in range(self.max_threads)])
|
2022-08-25 09:46:18 +00:00
|
|
|
self.eval_callback = EvalCallback(self.eval_env, deterministic=True,
|
2022-08-25 10:29:48 +00:00
|
|
|
render=False, eval_freq=len(train_df),
|
2022-09-23 17:17:27 +00:00
|
|
|
best_model_save_path=str(dk.data_path))
|