diff --git a/doc/tutorials/reinforcement_learning/reinforcement_learning.txt b/doc/tutorials/reinforcement_learning/reinforcement_learning.txt index ab712e98ef..d8fb6cec58 100644 --- a/doc/tutorials/reinforcement_learning/reinforcement_learning.txt +++ b/doc/tutorials/reinforcement_learning/reinforcement_learning.txt @@ -48,9 +48,10 @@ which are used by the agents while learning. - \c State: The State class is a representation of the environment. For the CartPole, this would involve storing the position, velocity, angle and angular velocity. -- \c Action: It is an enum naming all the possible actions the agent can take in the - environment. Continuing with the CartPole example, the Action enum would simply - contain the two possible actions, backward and forward. +- \c Action: For discrete environments, Action is a class with an enum naming all the possible + actions the agent can take in the environment. Continuing with the CartPole example, the enum + would simply contain the two possible actions, `backward` and `forward`. For continuous environments, + the Action class contains an array with its size depending on the action space. - \c Sample: This method is perhaps the heart of the environment, providing rewards to the agent depending on the state and the action taken, and updates the state based on @@ -123,6 +124,7 @@ agent on the CartPole environment. The code has been broken into chunks for easy #include #include #include +#include #include #include #include @@ -136,6 +138,23 @@ using namespace mlpack::rl; We include all the necessary components of our toy example and declare namespaces for convenience. +@code +int main() +{ + // Set up the network. + SimpleDQN<> model(4, 64, 32, 2); +@endcode + +The first step in setting our Q-learning agent is to setup the network for it to use. SimpleDQN class creates a +simple feed forward network with 2 hidden layers. The network constructed here has an input shape of 4 and +output shape of 2. This corresponds to the structure of the CartPole environment, where each state is +represented as a column vector with 4 data members (position, velocity, angle, angular velocity). Similarly, +the output shape is represented by the number of possible actions, which in this case, is only 2 +(`foward` and `backward`). + +We can also use mlpack's ann module to setup a custom FFN network. For example, here we use a single +hidden layer. + @code int main() { @@ -150,15 +169,6 @@ int main() @endcode -The first step in setting our Q-learning agent is to setup the network for it to use. Here, -we use mlpack's ann module to setup a simple FFN network, consisting of a single hidden layer. - -@note -The network constructed here has an input shape of 4 and output shape of 2. This corresponds to -the structure of the CartPole environment, where each state is represented as a column vector with -4 data members (position, velocity, angle, angular velocity). Similarly, the output shape is represented -by the number of possible actions, which in this case, is only 2 (foward and backward). - The next step would be to setup the other components of the Q-learning agent, namely its policy, replay method and hyperparameters. @@ -180,18 +190,16 @@ And now, we get to the heart of the program, declaring a Q-Learning agent. @code QLearning - agent(std::move(config), std::move(model), std::move(policy), - std::move(replayMethod)); + agent(config, model, policy, replayMethod); @endcode Here, we call the `QLearning` constructor, passing in the type of environment, network, updater, policy and replay. We use `decltype(var)` as a shorthand for the variable, saving us the trouble of copying the lengthy templated type. -Similarly, `std::move` is called for convenience, moving the components instead of -duplicating them and copying them over. +We pass references of the objects we created, as parameters to QLearning class. -We have our Q-Learning agent `agent` ready to be trained on the Cart Pole environment. +Now, we have our Q-Learning agent `agent` ready to be trained on the Cart Pole environment. @code arma::running_stat averageReturn;