DQN cu experience replay

Deep Reinforcement Learning în Python

Timothée Carayol

Principal Machine Learning Engineer, Komment

Introducere în experience replay

 

  • Agentul DQN de bază învață doar din ultima experiență
    • Actualizările consecutive sunt puternic corelate
    • Agentul uită rapid
  • Soluție: Experience Replay
    • Stocarea experiențelor într-un buffer
    • La fiecare pas, învățare dintr-un batch aleatoriu de experiențe trecute

 

O fotografie aeriană a unui labirint de garduri vii

Deep Reinforcement Learning în Python

Coada cu două capete

 

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])
  • Peste capacitate, elementele cele mai vechi sunt eliminate

Un deque cu capacitatea șapte, conținând patru elemente etichetate de la unu la patru.

Un deque cu capacitatea șapte, conținând patru elemente etichetate de la unu la patru. Patru elemente suplimentare etichetate de la cinci la opt sunt vizibile în dreapta.

Un deque cu capacitatea șapte, conținând șapte elemente etichetate de la unu la șapte. Un element suplimentar etichetat 8 este vizibil în dreapta.

Un deque cu capacitatea șapte, conținând șapte elemente etichetate de la doi la opt. Un element suplimentar etichetat 1 este vizibil în stânga.

Deep Reinforcement Learning în Python

Implementarea 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)
...
  • Memorie replay: deque cu capacitate limitată
  • .push():
    • Experiența ca tuplu de tranziție
    • Adăugare experiență în buffer
    • La capacitate maximă: elimină cea mai veche experiență
Deep Reinforcement Learning în Python

Implementarea 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

 

  • Eșantionare aleatorie din experiențe anterioare
  • batch: din listă de tupluri de tranziție...
  • ...la tuplu de liste...
  • ...la tuplu de tensori PyTorch
Deep Reinforcement Learning în Python

Integrarea Experience Replay în DQN

  1. Înainte de bucla de antrenare: replay_buffer = ReplayBuffer(10000)

  2. În bucla de antrenare, după selectarea acțiunii:

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)
  • Inițializare replay buffer
  • Adăugare tranziție recentă în buffer

Dacă lungimea bufferului $\geq$ batch_size:

  • Extragere batch aleatoriu și calculul pierderii
  • Calculul pierderii conceptual nemodificat
  • Mean Squared Bellman Error pe un batch din memorie
    • Învățare mai stabilă și mai eficientă
Deep Reinforcement Learning în Python

Să exersăm!

Deep Reinforcement Learning în Python

Preparing Video For Download...