diff --git a/src/mlpack/methods/reinforcement_learning/sac_impl.hpp b/src/mlpack/methods/reinforcement_learning/sac_impl.hpp index 8da1e7f6a9..c36570aa31 100644 --- a/src/mlpack/methods/reinforcement_learning/sac_impl.hpp +++ b/src/mlpack/methods/reinforcement_learning/sac_impl.hpp @@ -170,9 +170,9 @@ void SAC< sampledNextStates); arma::rowvec Q1, Q2; targetQ1Network.Predict(targetQInput, Q1); - // targetQ2Network.Predict(targetQInput, Q2); + targetQ2Network.Predict(targetQInput, Q2); arma::rowvec nextQ = sampledRewards + config.Discount() * ((1 - isTerminal) - % Q1); + % arma::min(Q1, Q2)); arma::mat sampledActionValues(action.size, sampledActions.size()); for (size_t i = 0; i < sampledActions.size(); i++) @@ -180,11 +180,11 @@ void SAC< arma::mat learningQInput = arma::join_vert(sampledActionValues, sampledStates); learningQ1Network.Forward(learningQInput, Q1); - // learningQ2Network.Forward(learningQInput, Q2); + learningQ2Network.Forward(learningQInput, Q2); arma::mat gradQ1Loss, gradQ2Loss; lossFunction.Backward(Q1, nextQ, gradQ1Loss); - // lossFunction.Backward(Q2, nextQ, gradQ2Loss); + lossFunction.Backward(Q2, nextQ, gradQ2Loss); // Update the critic networks. arma::mat gradientQ1, gradientQ2; @@ -196,23 +196,27 @@ void SAC< qNetworkUpdatePolicy->Update(learningQ1Network.Parameters(), config.StepSize(), gradientQ1); #endif - // learningQ2Network.Backward(learningQInput, gradQ2Loss, gradientQ2); - // #if ENS_VERSION_MAJOR == 1 - // qNetworkUpdater.Update(learningQ2Network.Parameters(), config.StepSize(), - // gradientQ2); - // #else - // qNetworkUpdatePolicy->Update(learningQ2Network.Parameters(), - // config.StepSize(), gradientQ2); - // #endif + learningQ2Network.Backward(learningQInput, gradQ2Loss, gradientQ2); + #if ENS_VERSION_MAJOR == 1 + qNetworkUpdater.Update(learningQ2Network.Parameters(), config.StepSize(), + gradientQ2); + #else + qNetworkUpdatePolicy->Update(learningQ2Network.Parameters(), + config.StepSize(), gradientQ2); + #endif // Actor network update. - // arma::rowvec pi; - // policyNetwork.Predict(sampledStates, pi); + arma::rowvec pi; + policyNetwork.Predict(sampledStates, pi); - // arma::mat qInput = arma::join_vert(pi, sampledStates); - // learningQ1Network.Predict(qInput, Q1); - // learningQ2Network.Predict(qInput, Q2); + arma::mat qInput = arma::join_vert(pi, sampledStates); + learningQ1Network.Predict(qInput, Q1); + learningQ2Network.Predict(qInput, Q2); + + // Get the size of the first hidden layer in the Q network. + size_t hidden1 = boost::get *> + (learningQ1Network.Model()[0])->OutputSize(); arma::mat gradient; for (size_t i = 0; i < sampledStates.n_cols; i++) @@ -221,17 +225,26 @@ void SAC< arma::colvec singleState = sampledStates.col(i); arma::colvec singlePi; policyNetwork.Forward(singleState, singlePi); - arma::colvec input = arma::join_vert(singlePi, singleState); - learningQ1Network.Forward(input, q); - learningQ1Network.Backward(input, -1, gradQ); + arma::rowvec weightLastLayer; + + if (Q1(i) < Q2(i)) + { + learningQ1Network.Forward(input, q); + learningQ1Network.Backward(input, -1, gradQ); + weightLastLayer = learningQ1Network.Parameters(). + rows(0, hidden1 - 1).t(); + } + else + { + learningQ2Network.Forward(input, q); + learningQ2Network.Backward(input, -1, gradQ); + weightLastLayer = learningQ2Network.Parameters(). + rows(0, hidden1 - 1).t(); + } - size_t hidden1 = boost::get *> - (learningQ1Network.Model()[0])->OutputSize(); arma::colvec gradQBias = gradQ(input.n_rows * hidden1, 0, arma::size(hidden1, 1)); - arma::rowvec weightLastLayer = learningQ1Network.Parameters(). - rows(0, hidden1-1).t(); arma::mat gradPolicy = weightLastLayer * gradQBias; policyNetwork.Backward(singleState, gradPolicy, grad); if (i == 0)