From 31294089bf5657ea67de5fc6b40dcd90dbb2385d Mon Sep 17 00:00:00 2001 From: nishantkr18 Date: Mon, 25 May 2020 19:34:26 +0530 Subject: [PATCH] advantage and value networks work! --- src/mlpack/methods/ann/layer/concat.hpp | 8 ---- .../q_networks/dueling_dqn.hpp | 40 ++++++++++--------- 2 files changed, 21 insertions(+), 27 deletions(-) diff --git a/src/mlpack/methods/ann/layer/concat.hpp b/src/mlpack/methods/ann/layer/concat.hpp index 31f7b9a7e6..30bfe4f3b5 100644 --- a/src/mlpack/methods/ann/layer/concat.hpp +++ b/src/mlpack/methods/ann/layer/concat.hpp @@ -138,14 +138,6 @@ class Concat arma::Mat& gradient, const size_t index); - /* - * Add a new module to the model. - * - * @param layer The Layer to be added to the model. - */ - template - void Add(const LayerType& layer) { network.push_back(new LayerType(layer)); } - /* * Add a new module to the model. * 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 93f4c68d01..757ed33c76 100644 --- a/src/mlpack/methods/reinforcement_learning/q_networks/dueling_dqn.hpp +++ b/src/mlpack/methods/reinforcement_learning/q_networks/dueling_dqn.hpp @@ -29,8 +29,8 @@ using namespace mlpack::ann; */ template < typename FeatureNetworkType = FFN, GaussianInitialization>, - typename AdvantageNetworkType = Sequential<>, - typename ValueNetworkType = Sequential<> + typename AdvantageNetworkType = Sequential<>*, + typename ValueNetworkType = Sequential<>* > class DuelingDQN { @@ -57,21 +57,24 @@ class DuelingDQN valueNetwork(), advantageNetwork() { - valueNetwork.Add(new Linear<>(h1, h2)); - valueNetwork.Add(new ReLULayer<>()); - valueNetwork.Add(new Linear<>(h2, 1)); + valueNetwork = new Sequential<>(); + valueNetwork->Add(new Linear<>(h1, h2)); + valueNetwork->Add(new ReLULayer<>()); + valueNetwork->Add(new Linear<>(h2, 1)); - advantageNetwork.Add(new Linear<>(h1, h2)); - advantageNetwork.Add(new ReLULayer<>()); - advantageNetwork.Add(new Linear<>(h1, outputDim)); + advantageNetwork = new Sequential<>(); + advantageNetwork->Add(new Linear<>(h1, h2)); + advantageNetwork->Add(new ReLULayer<>()); + advantageNetwork->Add(new Linear<>(h2, outputDim)); - Concat<> concat = new Concat<>(); - concat.Add>(valueNetwork); - concat.Add>(advantageNetwork); + Concat<>* concat = new Concat<>(true); + concat->Add(valueNetwork); + concat->Add(advantageNetwork); featureNetwork.Add(new Linear<>(inputDim, h1)); featureNetwork.Add(new ReLULayer<>()); featureNetwork.Add(concat); + this->ResetParameters(); } DuelingDQN(FeatureNetworkType featureNetwork, @@ -81,10 +84,11 @@ class DuelingDQN advantageNetwork(std::move(advantageNetwork)), valueNetwork(std::move(valueNetwork)) { - Concat<> concat = new Concat<>(); - concat.Add>(valueNetwork); - concat.Add>(advantageNetwork); + Concat<>* concat = new Concat<>(true); + concat->Add(valueNetwork); + concat->Add(advantageNetwork); featureNetwork.Add(concat); + this->ResetParameters(); } /** @@ -114,12 +118,12 @@ class DuelingDQN */ void Forward(const arma::mat state, arma::mat& output) { - // arma::mat output, advantage, value; - // featureNetwork.Forward(state, output); + arma::mat advantage, value; + featureNetwork.Forward(state, output); // actionValue = advantage.each_row() + // (value - arma::mean(arma::mean(advantage))); - // networkOutput = output; + networkOutput = output; } /** @@ -151,8 +155,6 @@ class DuelingDQN void ResetParameters() { featureNetwork.ResetParameters(); - advantageNetwork.ResetParameters(); - valueNetwork.ResetParameters(); } //! Return the Parameters.