reverting to using two q networks as critic

This commit is contained in:
nishantkr18
2020-08-05 15:01:29 +05:30
parent 1f5fb8de2a
commit 5facccf10c
@@ -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<mlpack::ann::Linear<> *>
(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<mlpack::ann::Linear<> *>
(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)