Chuyển đến nội dung chính

Lesson 8: SAC & Advanced Algorithms — Off-Policy Deep RL

SAC: Soft Actor-Critic — maximum entropy RL. TD3: Twin Delayed DDPG. Continuous action spaces. DDPG. Model-based RL overview.

🧠 AI & ML — Lesson 7 Lesson 8: SAC & Advanced Algorithms — Off-Policy Deep Rs

Reinforcement Learning: From Basics to Advanced

Part 2: Deep Reinforcement Learning — Neural Networks meet RL

xdev.asia

Introduction

SAC (Soft Actor-Critic) is a state-of-the-art off-policy algorithm for continuous action spaces — combining maximum entropy objective with actor-critic. SAC automatically tunes temperature, is stable and sample-efficient.


1. Maximum Entropy RL

$$\pi^* = \arg\max_\pi \sum_t \mathbb{E}_{\pi}[r(s_t, a_t) + \alpha \mathcal{H}(\pi(\cdot|s_t))]$$

Key insight: Balance reward maximization AND entropy (exploration) → robust policies.


2. SAC Components

class SACActorContinuous(nn.Module):
    """Gaussian policy for continuous actions"""
    def __init__(self, state_dim, action_dim, hidden=256):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(state_dim, hidden), nn.ReLU(),
            nn.Linear(hidden, hidden), nn.ReLU(),
        )
        self.mean = nn.Linear(hidden, action_dim)
        self.log_std = nn.Linear(hidden, action_dim)
    
    def forward(self, state):
        feat = self.net(state)
        mean = self.mean(feat)
        log_std = self.log_std(feat).clamp(-20, 2)
        return mean, log_std
    
    def sample(self, state):
        mean, log_std = self.forward(state)
        std = log_std.exp()
        dist = torch.distributions.Normal(mean, std)
        x = dist.rsample()  # Reparameterization trick
        action = torch.tanh(x)
        log_prob = dist.log_prob(x) - torch.log(1 - action.pow(2) + 1e-6)
        return action, log_prob.sum(-1)

Twin Q-Networks

class TwinQNetwork(nn.Module):
    def __init__(self, state_dim, action_dim, hidden=256):
        super().__init__()
        self.q1 = nn.Sequential(
            nn.Linear(state_dim + action_dim, hidden), nn.ReLU(),
            nn.Linear(hidden, hidden), nn.ReLU(),
            nn.Linear(hidden, 1),
        )
        self.q2 = nn.Sequential(
            nn.Linear(state_dim + action_dim, hidden), nn.ReLU(),
            nn.Linear(hidden, hidden), nn.ReLU(),
            nn.Linear(hidden, 1),
        )
    
    def forward(self, state, action):
        x = torch.cat([state, action], dim=-1)
        return self.q1(x), self.q2(x)

3. TD3 — Twin Delayed DDPG

3 tricks on DDPG:

  1. Twin critics: Take minimum → reduce overestimation
  2. Delayed policy updates: Update actor less than critic
  3. Target policy smoothing: Add noise to target actions

4. Algorithm Selection Guide

AlgorithmAction SpaceOn/Off-PolicySample EfficiencyBest For
DQNDiscreteOff-policyGoodGames
PPOBothOn-policyLowGeneral purpose
SACContinuousOff-policyHighRobotics
TD3ContinuousOff-policyHighRobotics
DDPGContinuousOff-policyGoodSimple continuous

5. Model-based RL Overview

Instead of model-free, learn model of environment → plan:

ApproachLearn Model?Plan?Example
Model-freeNoNoDQN, PPO
Model-basedYesYesDreamer, MuZero
HybridYesYesWorld Models

Summary

ConceptsDescription
SACMax entropy + twin critics + auto temperature
TD3Twin critics + delayed + smoothing
EntropyEncourages exploration, robust policies
Off-policySample efficient, replay buffer