DQN s experience replay

Deep Reinforcement Learning v Pythonu

Timothée Carayol

Principal Machine Learning Engineer, Komment

Úvod do experience replay

 

  • Základní DQN agent se učí pouze z nejnovější zkušenosti
    • Po sobě jdoucí aktualizace jsou silně korelovány
    • Agent rychle zapomíná
  • Řešení: Experience Replay
    • Ukládání zkušeností do bufferu
    • V každém kroku se učí z náhodné dávky minulých zkušeností

 

Letecký snímek živého labyrintu

Deep Reinforcement Learning v Pythonu

Oboustranná fronta (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])
  • Po překročení kapacity jsou nejstarší prvky odstraněny

Deque s kapacitou sedm obsahující čtyři prvky označené jednou až čtyřmi.

Deque s kapacitou sedm obsahující čtyři prvky označené jednou až čtyřmi. Napravo jsou vidět čtyři další prvky označené pěti až osmi.

Deque s kapacitou sedm obsahující sedm prvků označených jednou až sedmi. Napravo je vidět jeden další prvek označený 8.

Deque s kapacitou sedm obsahující sedm prvků označených dvěma až osmi. Nalevo je vidět jeden další prvek označený 1.

Deep Reinforcement Learning v Pythonu

Implementace Replay Bufferu

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)
...
  • Paměť pro replay: deque s omezenou kapacitou
  • .push():
    • Zkušenost jako n-tice přechodu
    • Přidání zkušenosti do bufferu
    • Po dosažení kapacity: odstraní nejstarší zkušenost
Deep Reinforcement Learning v Pythonu

Implementace Replay Bufferu

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

 

  • Náhodný výběr z minulých zkušeností
  • batch: ze seznamu n-tic přechodů...
  • ...na n-tici seznamů...
  • ...na n-tici tensorů PyTorch
Deep Reinforcement Learning v Pythonu

Integrace Experience Replay do DQN

  1. Před trénovací smyčkou: replay_buffer = ReplayBuffer(10000)

  2. V trénovací smyčce, po výběru akce:

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)
  • Inicializace replay bufferu
  • Uložení posledního přechodu do bufferu

Pokud délka bufferu $\geq$ batch_size:

  • Náhodný výběr dávky z bufferu a výpočet ztráty
  • Výpočet ztráty zůstává konceptuálně nezměněn
  • Střední kvadratická Bellmanova chyba na dávce z paměti
    • Učení je stabilnější a efektivnější
Deep Reinforcement Learning v Pythonu

Pojďme si procvičit!

Deep Reinforcement Learning v Pythonu

Preparing Video For Download...