方策勾配におけるバッチ更新

Pythonで学ぶDeep Reinforcement Learning

Timothée Carayol

Principal Machine Learning Engineer, Komment

逐次更新 vs バッチ更新

エピソードを表す大きな枠。

Pythonで学ぶDeep Reinforcement Learning

逐次更新 vs バッチ更新

大きな枠の中にステップ1の小枠。「行動を選択」のボックス。

Pythonで学ぶDeep Reinforcement Learning

逐次更新 vs バッチ更新

ステップ1枠に「環境を反復」の小ボックスが追加。

Pythonで学ぶDeep Reinforcement Learning

逐次更新 vs バッチ更新

ステップ1の下に「損失を計算」「勾配降下」の枠。

Pythonで学ぶDeep Reinforcement Learning

逐次更新 vs バッチ更新

同一の枠がステップ2にも表示される。

Pythonで学ぶDeep Reinforcement Learning

逐次更新 vs バッチ更新

ステップ3とステップ4も表示される。

Pythonで学ぶDeep Reinforcement Learning

A2C / PPO の更新をバッチ化

大きなエピソード枠の半分に「ロールアウト1」。中に「ステップ1」「ステップ2」の空枠。

Pythonで学ぶDeep Reinforcement Learning

A2C / PPO の更新をバッチ化

ステップ1枠内に「行動を選択」「環境を反復」。

Pythonで学ぶDeep Reinforcement Learning

A2C / PPO の更新をバッチ化

ステップ2も同様。

Pythonで学ぶDeep Reinforcement Learning

A2C / PPO の更新をバッチ化

ステップ1と2の下に「損失を計算」「勾配降下」が1回ずつ表示。

Pythonで学ぶDeep Reinforcement Learning

A2C / PPO の更新をバッチ化

残り半分に同じ2ステップの「ロールアウト2」。

Pythonで学ぶDeep Reinforcement Learning

バッチ更新付きの A2C 学習ループ

 

# Set rollout length
rollout_length = 10

# Initiate loss batches
actor_losses = torch.tensor([]) critic_losses = torch.tensor([])
  • 損失のバッチを初期化
  • いつも通りエピソードとステップを反復

 

for episode in range(10):
  state, info = env.reset()
  done = False
  while not done:
    action, action_log_prob = select_action(actor, 
                                            state)                
    next_state, reward, terminated, truncated, _ = (
                                   env.step(action))
    done = terminated or truncated    
    actor_loss, critic_loss = calculate_losses(
        critic, action_log_prob, 
        reward, state, next_state, done)
    ...
Pythonで学ぶDeep Reinforcement Learning

バッチ更新付きの A2C 学習ループ

  ...
  actor_losses = torch.cat((actor_losses, actor_loss))
  critic_losses = torch.cat((critic_losses, critic_loss))

# If rollout is full, update the networks if len(actor_losses) >= rollout_length:
actor_loss_batch = actor_losses.mean() critic_loss_batch = critic_losses.mean()
actor_optimizer.zero_grad() actor_loss_batch.backward() actor_optimizer.step() critic_optimizer.zero_grad() critic_loss_batch.backward() critic_optimizer.step()
actor_losses = torch.tensor([]) critic_losses = torch.tensor([])
state = next_state

 

  • 各ステップの損失をバッチに追加
  • ロールアウトが満杯なら
    • .mean() でバッチ平均損失を計算
    • 勾配降下を実行
    • 損失バッチを再初期化
Pythonで学ぶDeep Reinforcement Learning

複数エージェントでの A2C / PPO

 

2本の横帯がエージェント1・2を表す。各エージェントは長さの異なる4・3エピソードを経験。各エピソード内にステップ枠。下に8ステップ幅のロールアウト枠が3つあり、それぞれ「損失を計算」「勾配降下」のラベル。上部凡例:「ロールアウト長: 8ステップ;エージェント数: 2」。

Pythonで学ぶDeep Reinforcement Learning

ロールアウトとミニバッチ

前スライド同様の2本の帯。下の3つのロールアウト枠は上に「シャッフル」欄、その下を縦に4分割し各「ミニバッチ」内に「損失を計算」「勾配降下」。凡例:「ロールアウト長: 8ステップ;ミニバッチサイズ: 4 (2x2);エージェント数: 2」。

Pythonで学ぶDeep Reinforcement Learning

複数エポックの PPO

前図に似るが、各ロールアウト枠が縦にも4分割。上から「シャッフル」、「エポック1」(内に4つのミニバッチ)、「再シャッフル」、「エポック2」(同様)。凡例:「ロールアウト長: 8ステップ;ミニバッチサイズ: 4 (2x2);エージェント数: 2;エポック数: 2」。

Pythonで学ぶDeep Reinforcement Learning

練習しましょう!

Pythonで学ぶDeep Reinforcement Learning

Preparing Video For Download...