Deep Deterministic Policy Gradient (DDPG) Implementation in PyTorch
This implementation builds Deep Deterministic Policy Gradient (DDPG) from scratch. It includes a tanh-bounded deterministic Actor, a Q-value Critic over state-action pairs, a replay buffer that samples mini-batches of transitions, target networks updated via soft Polyak averaging, and a smoke test that exercises action selection and a short training loop.
import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim
from collections import deque
import random
class Actor(nn.Module):
def __init__(self, state_dim, action_dim, hidden_dim=256):
super(Actor, self).__init__()
self.network = nn.Sequential(
nn.Linear(state_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, action_dim),
nn.Tanh() # Bound actions to [-1, 1]
)
def forward(self, state):
return self.network(state)
class Critic(nn.Module):
def __init__(self, state_dim, action_dim, hidden_dim=256):
super(Critic, self).__init__()
self.network = nn.Sequential(
nn.Linear(state_dim + action_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, 1)
)
def forward(self, state, action):
x = torch.cat([state, action], dim=1)
return self.network(x)
class ReplayBuffer:
def __init__(self, capacity):
self.buffer = deque(maxlen=capacity)
def push(self, state, action, reward, next_state, done):
self.buffer.append((state, action, reward, next_state, done))
def sample(self, batch_size):
transitions = random.sample(self.buffer, batch_size)
state, action, reward, next_state, done = zip(*transitions)
return (
torch.FloatTensor(state),
torch.FloatTensor(action),
torch.FloatTensor(reward),
torch.FloatTensor(next_state),
torch.FloatTensor(done)
)
def __len__(self):
return len(self.buffer)
class DDPG:
def __init__(self, state_dim, action_dim, hidden_dim=256, buffer_size=1000000,
batch_size=64, gamma=0.99, tau=0.001, actor_lr=1e-4, critic_lr=1e-3):
# Initialize networks
self.actor = Actor(state_dim, action_dim, hidden_dim)
self.actor_target = Actor(state_dim, action_dim, hidden_dim)
self.actor_target.load_state_dict(self.actor.state_dict())
self.critic = Critic(state_dim, action_dim, hidden_dim)
self.critic_target = Critic(state_dim, action_dim, hidden_dim)
self.critic_target.load_state_dict(self.critic.state_dict())
self.actor_optimizer = optim.Adam(self.actor.parameters(), lr=actor_lr)
self.critic_optimizer = optim.Adam(self.critic.parameters(), lr=critic_lr)
self.replay_buffer = ReplayBuffer(buffer_size)
self.batch_size = batch_size
self.gamma = gamma # Discount factor
self.tau = tau # Soft update parameter
def select_action(self, state, noise_std=0.1):
with torch.no_grad():
state = torch.FloatTensor(state).unsqueeze(0)
action = self.actor(state).squeeze(0).numpy()
# Add exploration noise
action += np.random.normal(0, noise_std, size=action.shape)
return np.clip(action, -1, 1)
def train(self):
if len(self.replay_buffer) < self.batch_size:
return
state, action, reward, next_state, done = self.replay_buffer.sample(self.batch_size)
# Update critic
with torch.no_grad():
next_action = self.actor_target(next_state)
target_Q = self.critic_target(next_state, next_action)
target_Q = reward.unsqueeze(1) + (1 - done.unsqueeze(1)) * self.gamma * target_Q
current_Q = self.critic(state, action)
critic_loss = nn.MSELoss()(current_Q, target_Q)
self.critic_optimizer.zero_grad()
critic_loss.backward()
self.critic_optimizer.step()
# Update actor
actor_loss = -self.critic(state, self.actor(state)).mean()
self.actor_optimizer.zero_grad()
actor_loss.backward()
self.actor_optimizer.step()
# Soft update target networks
self._soft_update(self.actor_target, self.actor)
self._soft_update(self.critic_target, self.critic)
def _soft_update(self, target, source):
# θ_target = τ*θ_source + (1 - τ)*θ_target
for target_param, param in zip(target.parameters(), source.parameters()):
target_param.data.copy_(
self.tau * param.data + (1.0 - self.tau) * target_param.data
)
def test_ddpg():
state_dim = 3
action_dim = 2
agent = DDPG(state_dim, action_dim)
state = np.random.randn(state_dim)
action = agent.select_action(state)
assert action.shape == (action_dim,)
assert np.all(action >= -1) and np.all(action <= 1)
# Training loop
for _ in range(100):
state = np.random.randn(state_dim)
action = agent.select_action(state)
reward = np.random.rand()
next_state = np.random.randn(state_dim)
done = False
agent.replay_buffer.push(state, action, reward, next_state, done)
agent.train()
print("All tests passed!")
if __name__ == "__main__":
test_ddpg()