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

レッスン 8: SAC と高度なアルゴリズム — オフポリシーのディープ RL

SAC: Soft Actor-Critic — 最大エントロピー RL。 TD3: ツインディレイ DDPG。継続的なアクションスペース。 DDPG。モデルベースの RL の概要。

🧠 AI と ML — レッスン 7 レッスン 8: SAC と高度なアルゴリズム — ポリシー外のディープ R

強化学習: 基礎から高度まで

パート 2: 深層強化学習 — ニューラル ネットワークと RL の出会い

xdev.asia

はじめに

SAC (ソフト アクター-クリティック) は、最大エントロピー目標とアクター-クリティックを組み合わせた、連続アクション スペース用の最先端のオフポリシー アルゴリズムです。 SAC は温度を自動的に調整し、安定しており、サンプル効率が優れています。


1. 最大エントロピー RL

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

重要な洞察: 報酬の最大化とエントロピー (探索) のバランス → 堅牢なポリシー。


2. SAC コンポーネント

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)

ツイン Q ネットワーク

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 — ツイン遅延 DDPG

DDPG の 3 つのトリック:

  1. 双子の批判者: 最小限にする → 過大評価を減らす
  2. ポリシー更新の遅延: 更新アクターが批判者よりも少ない
  3. ターゲット ポリシーのスムージング: ターゲット アクションにノイズを追加します。

4. アルゴリズムの選択ガイド

アルゴリズムアクションスペースオン/オフポリシーサンプル効率最適な用途
DQN離散ポリシー外良いゲーム
PPO両方オンポリシー低い汎用
SAC連続ポリシー外高ロボット工学
TD3連続ポリシー外高ロボット工学
DDPG連続ポリシー外良い単純な連続

5. モデルベースの RL の概要

モデルフリーの代わりに、環境のモデルを学習 → 計画します。

アプローチモデルを学習しますか?プラン?例
モデルフリーいいえいいえDQN、PPO
モデルベースはいはいドリーマー、ミューゼロ
ハイブリッドはいはいワールドモデル

概要

コンセプト説明
SAC最大エントロピー + 双子の批評家 + 自動温度
TD3ツイン批評家 + 遅延 + 平滑化
エントロピー探求を奨励し、強力な政策を講じます。
ポリシー外効率的なサンプル再生バッファ