DQN med erfarenhetsåteruppspelning

Djup förstärkningsinlärning i Python

Timothée Carayol

Principal Machine Learning Engineer, Komment

Introduktion till erfarenhetsåteruppspelning

 

  • En grundläggande DQN-agent lär sig enbart från den senaste erfarenheten
    • Konsekutiva uppdateringar är starkt korrelerade
    • Agenten glömmer snabbt
  • Lösning: erfarenhetsåteruppspelning
    • Lagra erfarenheter i en buffert
    • Lär dig vid varje steg från ett slumpmässigt urval av tidigare erfarenheter

 

Flygbild över en häcklabyrInt

Djup förstärkningsinlärning i Python

Den dubbeländade kön

 

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])
  • När kapaciteten överskrids tas de äldsta elementen bort

En deque med kapacitet sju, innehållande fyra element märkta ett till fyra.

En deque med kapacitet sju, innehållande fyra element märkta ett till fyra. Fyra ytterligare element märkta fem till åtta syns till höger.

En deque med kapacitet sju, innehållande sju element märkta ett till sju. Ett ytterligare element märkt 8 syns till höger.

En deque med kapacitet sju, innehållande sju element märkta två till åtta. Ett ytterligare element märkt 1 syns till vänster.

Djup förstärkningsinlärning i Python

Implementera återuppspelningsbuffert

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)
...
  • Återuppspelningsminne: deque med begränsad kapacitet
  • .push():
    • Erfarenhet som övergångstupel
    • Lägg till erfarenhet i bufferten
    • Vid full kapacitet: äldsta erfarenheten tas bort
Djup förstärkningsinlärning i Python

Implementera återuppspelningsbuffert

...
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

 

  • Slumpmässigt urval från tidigare erfarenheter
  • batch: från lista av övergångstuplar...
  • ...till tupel av listor...
  • ...till tupel av PyTorch-tensorer
Djup förstärkningsinlärning i Python

Integrera erfarenhetsåteruppspelning i DQN

  1. Före träningsloopen: replay_buffer = ReplayBuffer(10000)

  2. I träningsloopen, efter aktionsval:

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)
  • Initiera återuppspelningsbufferten
  • Lägg till senaste övergången i bufferten

Om buffertlängden $\geq$ batch_size:

  • Dra ett slumpmässigt urval från bufferten och beräkna förlusten
  • Förlustberäkningen är konceptuellt oförändrad
  • Mean Squared Bellman Error på ett urval från återuppspelningsminnet
    • Inlärningen blir stabilare och mer effektiv
Djup förstärkningsinlärning i Python

Nu kör vi en övning!

Djup förstärkningsinlärning i Python

Preparing Video For Download...