Deep Reinforcement Learning
While classical Q-Learning works for simple games like Tic-Tac-Toe, it fails when environments have infinite possible states (like driving a car or playing Mario).
Deep RL solves this by replacing the Q-Table with a Deep Neural Network.
Deep Q-Networks (DQN)
Instead of looking up a value in a table, we pass the current state (e.g., the pixels on a screen) into a CNN. The neural network outputs the predicted Q-Values for every possible action. This was the algorithm DeepMind used to beat human champions at Atari games.
Policy Gradients and PPO
DQN tries to predict the value of an action. Policy Gradients, on the other hand, try to learn the policy directly—the neural network outputs the probability of taking each action.
Proximal Policy Optimization (PPO): PPO is currently the most successful and widely used RL algorithm in the world. It is a policy gradient method that takes small, safe steps when updating the network to ensure the agent doesn't accidentally "forget" how to perform a task and collapse its learning progress.
[!IMPORTANT] PPO is the algorithm OpenAI used to train ChatGPT! The final step of building an LLM is RLHF (Reinforcement Learning from Human Feedback), where a PPO agent learns to generate text that maximizes human upvotes.
Python Libraries for Deep RL
Writing RL algorithms from scratch is notoriously difficult to debug. Standard practice is to use tested libraries:
- Stable-Baselines3: The industry standard for Deep RL algorithms built on [PyTorch](../../Course-1-Mathematics and Frameworks/Ch-9 Deep-Learning-Frameworks/PyTorch.mdx).
- Ray RLlib: For distributed, enterprise-scale reinforcement learning.
# Example: Training a PPO agent to play a gym environment using Stable-Baselines3
# !pip install stable-baselines3 gym
import gym
from stable_baselines3 import PPO
# 1. Create the environment
env = gym.make("CartPole-v1")
# 2. Initialize the PPO Agent
model = PPO("MlpPolicy", env, verbose=1)
# 3. Train the Agent!
model.learn(total_timesteps=10000)
# 4. Save the model
model.save("ppo_cartpole")