近接方策最適化 (PPO)

Pythonで学ぶDeep Reinforcement Learning

Timothée Carayol

Principal Machine Learning Engineer, Komment

A2C

  • A2C の方策更新:
    • 推定が不安定(揮発的)
    • 変化が大きく不安定になり得る
  • 性能を損なう可能性

荒地で事故に遭い故障した火星探査車

PPO

  • PPO は各方策更新量に上限を設ける
  • 安定性が向上

火星表面を順調に走行する探査車

Pythonで学ぶDeep Reinforcement Learning

確率比

  • PPO の主な革新点: 新しい目的関数
  • 中核は次の比率

確率比: 新しい方策下の行動確率と古い方策下の確率の比(r_t)

  • 行動 $a_t$ は $\theta$ で $\theta_{old}$ よりどれだけ起こりやすいか?

 

  ratio = action_log_prob.exp() / 
          old_action_log_prob.exp().detach()

# あるいは同等に ratio = torch.exp(action_log_prob - old_action_log_prob.detach())
  • 逆伝播を防ぐため分母を detach する
Pythonで学ぶDeep Reinforcement Learning

確率比のクリッピング

 

  • クリップ関数:

clip(x, 0.8, 1.2) のグラフ。x=0.6〜1.4、x<0.8 で 0.8、0.8〜1.2 で x、x>1.2 で 1.2

クリップ済み確率比は clip(r_t, 1-epsilon, 1+epsilon)

 

 

clipped_ratio = torch.clamp(ratio,
                            1-epsilon, 
                            1+epsilon)
Pythonで学ぶDeep Reinforcement Learning

calculate_ratios 関数

 

def calculate_ratios(action_log_prob, action_log_prob_old, epsilon):

prob = action_log_prob.exp() prob_old = action_log_prob_old.exp() prob_old_detached = prob_old.detach() ratio = prob / prob_old_detached clipped_ratio = torch.clamp(ratio, 1-epsilon, 1+epsilon)
return (ratio, clipped_ratio)
epsilon = .2 の例:

Ratio: tensor(1.25)
Clipped ratio: tensor(1.20)
Pythonで学ぶDeep Reinforcement Learning

PPO の目的関数

 

J surr = E_t(r_t * advantage)

surr1 = ratio * td_error.detach()

surr2 = clipped_ratio * td_error.detach()
objective = torch.min(surr1, surr2)

 

  • クリップ付き代替目的:

$$\mathrm{clip}(r_t(\theta),1-\varepsilon,1+\varepsilon)\hat{A}$$

  • PPO のクリップ付き代替目的関数:

クリップ付き代替目的関数:ratio×advantage と clipped ratio×advantage の最小値の期待値

  • A2C より安定
Pythonで学ぶDeep Reinforcement Learning

PPO の損失計算

 

def calculate_losses(critic_network, 
                     action_log_prob,                                        
                     action_log_prob_old,
                     reward, state, next_state,
                     done
                     ):

    # TD 誤差を計算(A2C と同じ)
    value = critic_network(state)
    next_value = critic_network(next_state)
    td_target = (reward + 
                 gamma * next_value * (1-done))
    td_error = td_target - value
    ...

 

    ...
    ratio, clipped_ratio = 
            calculate_ratios(action_log_prob, 
                             action_log_prob_old,
                             epsilon)

surr1 = ratio * td_error.detach()
surr2 = clipped_ratio * td_error.detach()
objective = torch.min(surr1, surr2)
actor_loss = -objective
critic_loss = td_error ** 2 return actor_loss, critic_loss
Pythonで学ぶDeep Reinforcement Learning

Ayo berlatih!

Pythonで学ぶDeep Reinforcement Learning

Preparing Video For Download...