DQN avec relecture d'expériences

Deep Reinforcement Learning en Python

Timothée Carayol

Principal Machine Learning Engineer, Komment

Introduction à la relecture d'expériences

 

  • Un agent DQN minimal apprend seulement de la dernière expérience
    • Mises à jour consécutives fortement corrélées
    • L'agent oublie vite
  • Solution : relecture d'expériences
    • Stocker les expériences dans un tampon
    • À chaque pas, apprendre d'un lot aléatoire d'expériences passées

 

Vue aérienne d'un labyrinthe de haies

Deep Reinforcement Learning en Python

La file à double extrémité (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])
  • Au-delà de la capacité, les plus anciens éléments sont retirés

Une deque de capacité sept, contenant quatre éléments nommés un à quatre.

Une deque de capacité sept, contenant quatre éléments nommés un à quatre. Quatre éléments additionnels nommés cinq à huit sont visibles à droite.

Une deque de capacité sept, contenant sept éléments nommés un à sept. Un élément additionnel nommé 8 est visible à droite.

Une deque de capacité sept, contenant sept éléments nommés deux à huit. Un élément additionnel nommé 1 est visible à gauche.

Deep Reinforcement Learning en Python

Implémenter le tampon de relecture

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)
...
  • Mémoire de relecture : deque à capacité limitée
  • .push() :
    • Expérience comme transition en tuple
    • Ajoute l'expérience au tampon
    • À capacité atteinte : supprime l'expérience la plus ancienne
Deep Reinforcement Learning en Python

Implémenter le tampon de relecture

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

 

  • Tirer au hasard des expériences passées
  • batch : d'une liste de tuples de transition...
  • ...à un tuple de listes...
  • ...à un tuple de tenseurs PyTorch
Deep Reinforcement Learning en Python

Intégrer la relecture d'expériences au DQN

  1. Avant la boucle d'entraînement : replay_buffer = ReplayBuffer(10000)

  2. Dans la boucle, après la sélection de l'action :

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)
  • Initialiser le tampon de relecture
  • Envoyer la dernière transition dans le tampon

Si la taille du tampon $\geq$ batch_size :

  • Tirer un lot aléatoire du tampon et calculer la perte
  • Calcul de la perte inchangé conceptuellement
  • Erreur quadratique moyenne de Bellman sur un lot de relecture
    • Apprentissage plus stable et plus efficace
Deep Reinforcement Learning en Python

Passons à la pratique !

Deep Reinforcement Learning en Python

Preparing Video For Download...