Reinforcement Learning with Neural Networks: Mathematical Details
Watch on YouTube →
Overview
Josh Starmer of StatQuest details the mathematical underpinnings of training a neural network using reinforcement learning via the policy gradients method. He walks through calculating the derivative of cross-entropy with respect to a bias term, using the chain rule to combine derivatives of the sigmoid activation function and the cross-entropy loss, and then updating the bias based on a reward signal derived from the outcome of an action.
Key takeaways
- Reinforcement learning trains neural networks by adjusting parameters based on rewards received for actions, rather than direct error signals.
- The policy gradients method uses the chain rule to compute the gradient of the expected reward with respect to network parameters.
- The derivative of the cross-entropy loss with respect to the bias is calculated by composing derivatives of the loss, the sigmoid output, and the pre-activation input.
- A reward signal (positive for good outcomes, negative for bad) is used to scale the calculated derivative, effectively flipping its direction if the initial guess was wrong.
- Gradient descent then uses this reward-scaled derivative to update the bias, moving it towards values that increase the probability of desired actions.
- The training process involves repeated cycles of action selection, reward determination, derivative calculation, and parameter updates until convergence.
Chapters
- This StatQuest focuses on the mathematical details of reinforcement learning with neural networks.
- Assumes familiarity with gradient descent and basic RL concepts for neural networks.
- Uses an example of choosing between Squatch's Fry Shack and Norm's Fry Hut based on hunger level.
- An input of 0.0 (not hungry) is fed into the neural network.
- The network outputs probabilities P(Norm) = 0.5 and P(Squatch) = 0.5.
- A random number (0.2) is chosen, leading to a decision to visit Squatch's Fry Shack.
- Cross-entropy quantifies the difference between the ideal probability (1.0 for the chosen action) and the network's output.
- The derivative of cross-entropy with respect to the bias is calculated using the chain rule.
- This involves derivatives of cross-entropy w.r.t. P(Norm), P(Norm) w.r.t. X (sigmoid input), and X w.r.t. bias.
- The derivative of the sigmoid activation function is derived as sigmoid(x) * (1 - sigmoid(x)).
- The derivative of X (hunger * weight + bias) with respect to the bias is 1.
- The combined derivative of cross-entropy w.r.t. bias is calculated for both Squatch and Norm scenarios.
- The calculated derivative is multiplied by a reward (1.0 for correct guess, -1.0 for incorrect).
- This 'updated derivative' corrects the direction of the gradient descent step.
- Gradient descent uses the updated derivative and learning rate (1.0) to calculate a step size and update the bias.
- The process iterates with new actions and rewards, eventually converging the bias to approximately -10.
Summary, takeaways, and chapters were generated by AI from the video's transcript and may contain errors. The video belongs to its creator, StatQuest with Josh Starmer.