経験再生付きDQN

Pythonで学ぶDeep Reinforcement Learning

Timothée Carayol

Principal Machine Learning Engineer, Komment

経験再生の導入

 

  • 最小構成のDQNは直近の経験のみから学習
    • 連続更新は強く相関
    • 忘れやすい
  • 解決策: 経験再生
    • 経験をバッファに保存
    • 各ステップで過去経験のランダムバッチから学習

 

生け垣迷路の航空写真

Pythonで学ぶDeep Reinforcement Learning

両端キュー(deque)

 

from collections import deque

# Instantiate with limited capacity buffer = deque([1,2,3,4], maxlen=7)
# Extend to the right side buffer.extend([5,6,7,8])
  • 容量超過時は最古の要素が捨てられる

容量7のdequeに、1〜4の4要素が入っている図。

容量7のdequeに1〜4の4要素。右側に5〜8の追加要素が見える図。

容量7のdequeに1〜7の7要素が入っている図。右に8が見える。

容量7のdequeに2〜8の7要素が入っている図。左に1が見える。

Pythonで学ぶDeep Reinforcement Learning

Replay Bufferの実装

import random

class ReplayBuffer:
def __init__(self, capacity):
self.memory = deque([], maxlen=capacity)
def push(self, state, action, reward, next_state, done):
experience_tuple = (state, action, reward, next_state, done)
self.memory.append(experience_tuple)
def __len__(self): return len(self.memory)
...
  • リプレイメモリ: 容量制限付きdeque
  • .push():
    • 遷移タプルとして保存
    • バッファに追加
    • 容量到達で最古を破棄
Pythonで学ぶDeep Reinforcement Learning

Replay Bufferの実装

...
def sample(self, batch_size):

batch = random.sample(self.memory, batch_size)
states, actions, rewards, next_states, dones = ( zip(*batch))
states_tensor = torch.tensor( states, dtype=torch.float32) ... # repeat identically for # rewards, next_states, dones
actions_tensor = torch.tensor( actions, dtype=torch.long).unsqueeze(1)
return states_tensor, actions_tensor, rewards_tensor, next_states_tensor, dones_tensor

 

  • 過去の経験からランダム抽出
  • batch: 遷移タプルのリストから…
  • …リストのタプルへ…
  • …PyTorchテンソルのタプルへ
Pythonで学ぶDeep Reinforcement Learning

DQNへの経験再生の統合

  1. 学習ループ前: replay_buffer = ReplayBuffer(10000)

  2. 学習ループ内、行動選択の後:

replay_buffer.push((state, action, 
                    reward, next_state, done))

if len(replay_buffer) >= batch_size:
states, actions, rewards, next_states, dones = ( replay_buffer.sample(batch_size))
q_values = ( q_network(states).gather(1, actions).squeeze(1))
next_states_q_values = q_network(next_states).amax(1)
target_q_values = ( rewards + gamma * next_states_q_values * (1-dones))
loss = nn.MSELoss()(target_q_values, q_values)
  • リプレイバッファを初期化
  • 直近の遷移をバッファに追加

バッファ長が batch_size 以上なら:

  • バッファからランダムにバッチを取得し損失計算へ
  • 損失計算の概念は不変
  • リプレイメモリ上の平均二乗ベルマン誤差
    • 学習がより安定・効率的
Pythonで学ぶDeep Reinforcement Learning

Passons à la pratique !

Pythonで学ぶDeep Reinforcement Learning

Preparing Video For Download...