From de329c9c9fdf76fed6e4f2927a67f485ae0fdb34 Mon Sep 17 00:00:00 2001 From: Ryan Curtin Date: Wed, 15 Jun 2022 22:58:26 -0400 Subject: [PATCH] Adapt Q-learning network types to new API. --- .../reinforcement_learning/q_learning_impl.hpp | 11 +++++------ .../q_networks/categorical_dqn.hpp | 4 ++-- .../reinforcement_learning/q_networks/dueling_dqn.hpp | 9 +++------ .../reinforcement_learning/q_networks/simple_dqn.hpp | 4 ++-- 4 files changed, 12 insertions(+), 16 deletions(-) diff --git a/src/mlpack/methods/reinforcement_learning/q_learning_impl.hpp b/src/mlpack/methods/reinforcement_learning/q_learning_impl.hpp index 7ff52cde88..afd4bafe0c 100644 --- a/src/mlpack/methods/reinforcement_learning/q_learning_impl.hpp +++ b/src/mlpack/methods/reinforcement_learning/q_learning_impl.hpp @@ -52,10 +52,12 @@ QLearning< targetNetwork = learningNetwork; // Set up q-learning network. - if (learningNetwork.Parameters().is_empty()) - learningNetwork.Reset(); + if (learningNetwork.Parameters().n_elem != environment.InitialSample().Encode().n_elem) + learningNetwork.Reset(environment.InitialSample().Encode().n_elem); - targetNetwork.Reset(); + // Initialize the target network with the parameters of learning network. + targetNetwork.Parameters() = learningNetwork.Parameters(); + targetNetwork.Reset(environment.InitialSample().Encode().n_elem); #if ENS_VERSION_MAJOR == 1 this->updater.Initialize(learningNetwork.Parameters().n_rows, @@ -66,9 +68,6 @@ QLearning< learningNetwork.Parameters().n_rows, learningNetwork.Parameters().n_cols); #endif - - // Initialize the target network with the parameters of learning network. - targetNetwork.Parameters() = learningNetwork.Parameters(); } template < diff --git a/src/mlpack/methods/reinforcement_learning/q_networks/categorical_dqn.hpp b/src/mlpack/methods/reinforcement_learning/q_networks/categorical_dqn.hpp index 1b9ed9f5da..cd2313dbcd 100644 --- a/src/mlpack/methods/reinforcement_learning/q_networks/categorical_dqn.hpp +++ b/src/mlpack/methods/reinforcement_learning/q_networks/categorical_dqn.hpp @@ -167,9 +167,9 @@ class CategoricalDQN /** * Resets the parameters of the network. */ - void Reset() + void Reset(const size_t inputDimensionality = 0) { - network.Reset(); + network.Reset(inputDimensionality); } /** diff --git a/src/mlpack/methods/reinforcement_learning/q_networks/dueling_dqn.hpp b/src/mlpack/methods/reinforcement_learning/q_networks/dueling_dqn.hpp index 9358333119..d0a6020db5 100644 --- a/src/mlpack/methods/reinforcement_learning/q_networks/dueling_dqn.hpp +++ b/src/mlpack/methods/reinforcement_learning/q_networks/dueling_dqn.hpp @@ -126,7 +126,6 @@ class DuelingDQN completeNetwork.Add(featureNetwork); completeNetwork.Add(concat); - this->Reset(); } /** @@ -151,7 +150,6 @@ class DuelingDQN concat->Add(advantageNetwork); completeNetwork.Add(featureNetwork); completeNetwork.Add(concat); - this->Reset(); } //! Copy constructor. @@ -185,8 +183,7 @@ class DuelingDQN completeNetwork.Predict(state, networkOutput); value = networkOutput.row(0); advantage = networkOutput.rows(1, networkOutput.n_rows - 1); - actionValue = advantage.each_row() + - (value - arma::mean(advantage)); + actionValue = advantage.each_row() + (value - arma::mean(advantage)); } /** @@ -228,9 +225,9 @@ class DuelingDQN /** * Resets the parameters of the network. */ - void Reset() + void Reset(const size_t inputDimensionality = 0) { - completeNetwork.Reset(); + completeNetwork.Reset(inputDimensionality); } /** diff --git a/src/mlpack/methods/reinforcement_learning/q_networks/simple_dqn.hpp b/src/mlpack/methods/reinforcement_learning/q_networks/simple_dqn.hpp index dca87fa822..033aa3a9ab 100644 --- a/src/mlpack/methods/reinforcement_learning/q_networks/simple_dqn.hpp +++ b/src/mlpack/methods/reinforcement_learning/q_networks/simple_dqn.hpp @@ -118,9 +118,9 @@ class SimpleDQN /** * Resets the parameters of the network. */ - void Reset() + void Reset(const size_t inputDimensionality = 0) { - network.Reset(); + network.Reset(inputDimensionality); } /**