完全な DQN アルゴリズム

Pythonで学ぶDeep Reinforcement Learning

Timothée Carayol

Principal Machine Learning Engineer, Komment

DQN アルゴリズム

 

 

  • 経験再生つき DQN を学習
  • 初期の DQN(2015)に近い
  • まだ 2 つ不足:
    • エプシロングリーディー → 探索を増やす
    • 固定 Q ターゲット → 学習を安定化

冒険者が立ち、そばにギリシャ文字イプシロンがある

氷に閉じ込められた大文字の Q

Pythonで学ぶDeep Reinforcement Learning

DQN におけるエプシロングリーディー

  • select_action() に減衰エプシロングリーディーを実装
def select_action(q_values, step, start, end, decay):

# このステップのしきい値を計算 epsilon = ( end + (start-end) * math.exp(-step / decay))
# 0〜1 の一様乱数を生成 sample = random.random()
if sample < epsilon: # ランダムなアクション番号を返す return random.choice(range(len(q_values)))
# 最大 Q 値のアクション番号を返す return torch.argmax(q_values).item()
  • $\varepsilon = end + (start-end) \cdot e^{-\frac{step}{decay}}$
  • 確率 $\varepsilon$ でランダム行動
  • 確率 $1 - \varepsilon$ で最大値の行動

減衰パラメータごとのエプシロン減衰スケジュールのプロット。

Pythonで学ぶDeep Reinforcement Learning

固定 Q ターゲット

 

  • ベルマン誤差では:
    • Q 値と TD ターゲットの両方に同じ Q-Network を使用
    • ターゲットが変動し不安定化

 

  • ターゲットを安定させるため target network を導入

 

ベルマン誤差: (r_t+1 + gamma max(Q(s_t+1, a))) - Q(s_t, a_t)

(r_t+1 + gamma max(Q_target(s_t+1, a))) - Q_online(s_t, a_t)

Pythonで学ぶDeep Reinforcement Learning

固定 Q ターゲットの実装

online_network = QNetwork(state_size, action_size)
target_network = QNetwork(state_size, action_size)

target_network.load_state_dict( online_network.state_dict())
def update_target_network( target_network, online_network, tau):
target_net_state_dict = target_network.state_dict() online_net_state_dict = online_network.state_dict() for key in online_net_state_dict:
target_net_state_dict[key] = ( online_net_state_dict[key] * tau + target_net_state_dict[key] * (1 - tau))
target_network.load_state_dict( target_net_state_dict)
return None
  • 初期状態では Online Network = Target Network
  • ネットワークの state dict には全重みが入る: ネットワークの state dictionary の例。fc1.weight, fc1.bias, fc2.weight の各エントリがテンソル値を持つ。
  • 各ステップで Target Network の各重みを Online Network に少しずつ近づける
Pythonで学ぶDeep Reinforcement Learning

固定 Q ターゲットでの損失計算

# アクション選択後(内側ループ)
if len(replay_buffer) >= batch_size:
  states, actions, rewards, next_states, dones = 
      replay_buffer.sample(64)

q_values = (online_network(states) .gather(1, actions).squeeze(1))
with torch.no_grad():
next_q_values = ( target_network(next_states).amax(1)) target_q_values = ( rewards + gamma * next_q_values * (1 - dones))
loss = torch.nn.MSELoss()(target_q_values, q_values) optimizer.zero_grad() loss.backward() optimizer.step()
update_target_network( target_network, online_network, tau)

 

  • Q 値は online_network を使用
  • ターゲット Q 値は target_network を使用
  • ターゲット Q 値計算では torch.no_grad() で勾配追跡を無効化
  • 損失は引き続き平均二乗ベルマン誤差
  • update_target_network()target_network をゆっくり更新
Pythonで学ぶDeep Reinforcement Learning

Passons à la pratique !

Pythonで学ぶDeep Reinforcement Learning

Preparing Video For Download...