Double DQN

Djup förstärkningsinlärning i Python

Timothée Carayol

Principal Machine Learning Engineer, Komment

Double Q-learning

  • Q-learning överskattar Q-värden, vilket försämrar inlärningseffektiviteten
  • Detta beror på maximeringsbias
  • Double Q-Learning eliminerar bias genom att separera aktionsval och värdeestimering

Två Q-tabeller (illustration från kursen Reinforcement Learning with Gymnasium using Python); double Q-learning använder dem omväxlande

Djup förstärkningsinlärning i Python

Idén bakom DDQN

  • Utgå från ett komplett DQN (med fasta Q-mål)
  • I DQN TD-mål:
    • Aktionsval: målnätverk
    • Värdeestimering: målnätverk
  • I DDQN TD-mål:
    • Aktionsval: online-nätverk
    • Värdeestimering: målnätverk
  • Inte exakt double Q-learning (inga alternerande Q-nätverk)
  • Stor nytta med minimal förändring

Bellman Error (DQN with fixed Q-targets): Q_online(s_t, a_t) - (r_t+1 + gamma max(Q_target(s_t+1, a)))

Bellman Error (DDQN with fixed Q-targets): Q_online(s_t, a_t) - (r_t+1 + gamma Q_target(s_t+1, tilde a)) with tilde a = argmax_a(Q_online(s_t+1, a))

Djup förstärkningsinlärning i Python

Implementering av Double DQN

DQN:

... # instantiate online and target networks
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()(q_values, target_q_values) ... # gradient descent ... # target network update

DDQN:

... # instantiate online and target networks
q_values = (online_network(states)
            .gather(1, actions).squeeze(1))

with torch.no_grad():

target_q_values = (rewards + gamma * next_q_values * (1 - dones))
loss = torch.nn.MSELoss()(q_values, target_q_values) ... # gradient descent ... # target network update
Djup förstärkningsinlärning i Python

Implementering av Double DQN

DQN:

... # instantiate online and target networks
q_values = (online_network(states)
            .gather(1, actions).squeeze(1))

with torch.no_grad():
next_actions = (target_network(next_states) .argmax(1).unsqueeze(1))
next_q_values = (target_network(next_states) .gather(1, next_actions).squeeze(1))
target_q_values = (rewards + gamma * next_q_values * (1 - dones))
loss = torch.nn.MSELoss()(q_values, target_q_values) ... # gradient descent ... # target network update

DDQN:

... # instantiate online and target networks
q_values = (online_network(states)
            .gather(1, actions).squeeze(1))

with torch.no_grad():
next_actions = (online_network(next_states) .argmax(1).unsqueeze(1))
next_q_values = (target_network(next_states) .gather(1, next_actions).squeeze(1))
target_q_values = (rewards + gamma * next_q_values * (1 - dones))
loss = torch.nn.MSELoss()(q_values, target_q_values) ... # gradient descent ... # target network update
Djup förstärkningsinlärning i Python

DDQN-prestanda

 

  • Jämför prestanda för DDQN, DQN och mänskliga spelare på Atari-spel
  • DDQN: högre poäng än ursprungligt DQN
  • Gäller inte alltid – testa båda

Stapeldiagram som visar att DQN nästan matchar mänsklig prestanda för medianspelet och når övermänskliga poäng i genomsnitt, samt att DDQN slår både mänsklig prestanda och DQN på median och genomsnitt.

1 https://arxiv.org/abs/2303.11634
Djup förstärkningsinlärning i Python

Nu kör vi en övning!

Djup förstärkningsinlärning i Python

Preparing Video For Download...