Skip to content
AILinkDeepTech
Go back
Reinforcement Learning Advanced

Deep Deterministic Policy Gradient (DDPG) Implementation in PyTorch

Abstract

A PyTorch implementation of Deep Deterministic Policy Gradient (DDPG): an Actor that outputs tanh-bounded deterministic actions, a Q-value Critic over state-action pairs, a replay buffer, target networks with soft updates, and a smoke test that exercises action selection and a short training loop.

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()


Cite this Explanation

@article{ailinkdeeptech2025ddpgalgo,
  title={Deep Deterministic Policy Gradient (DDPG) Implementation in PyTorch},
  author={AILinkDeepTech},
  journal={AILinkDeepTech Algorithm Explanations},
  year={2025},
  url={https://ailinkdeeptech.com/research/ddpg_algo}
}

Related Explanations