reverting to using two q networks as critic
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user