DQN กับ Experience Replay

Deep Reinforcement Learning ด้วย Python

Timothée Carayol

Principal Machine Learning Engineer, Komment

แนะนำ Experience Replay

 

  • DQN พื้นฐานเรียนรู้จากประสบการณ์ล่าสุดเท่านั้น
    • การอัปเดตต่อเนื่องมีความสัมพันธ์กันสูง
    • Agent ลืมประสบการณ์เก่าได้ง่าย
  • วิธีแก้: Experience Replay
    • เก็บประสบการณ์ไว้ใน buffer
    • แต่ละขั้นตอน เรียนรู้จาก batch ของประสบการณ์ที่สุ่มจากอดีต

 

ภาพถ่ายทางอากาศของเขาวงกตพุ่มไม้

Deep Reinforcement Learning ด้วย Python

Double-Ended Queue

 

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])
  • เมื่อเกินความจุ รายการเก่าสุดจะถูกลบออก

deque ความจุ 7 ช่อง มี 4 สมาชิกที่ระบุหมายเลข 1 ถึง 4

deque ความจุ 7 ช่อง มี 4 สมาชิกที่ระบุหมายเลข 1 ถึง 4 พร้อมสมาชิกเพิ่มเติม 4 ตัวที่ระบุหมายเลข 5 ถึง 8 ทางด้านขวา

deque ความจุ 7 ช่อง มี 7 สมาชิกที่ระบุหมายเลข 1 ถึง 7 พร้อมสมาชิกเพิ่มเติม 1 ตัวที่ระบุหมายเลข 8 ทางด้านขวา

deque ความจุ 7 ช่อง มี 7 สมาชิกที่ระบุหมายเลข 2 ถึง 8 พร้อมสมาชิกเพิ่มเติม 1 ตัวที่ระบุหมายเลข 1 ทางด้านซ้าย

Deep Reinforcement Learning ด้วย Python

สร้าง 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)
...
  • Replay memory: deque ที่มีความจุจำกัด
  • .push():
    • เก็บประสบการณ์เป็น transition tuple
    • เพิ่มประสบการณ์เข้า buffer
    • เมื่อเต็มความจุ: ลบประสบการณ์เก่าสุดออก
Deep Reinforcement Learning ด้วย Python

สร้าง 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

 

  • สุ่มดึงข้อมูลจากประสบการณ์ในอดีต
  • batch: จาก list ของ transition tuple...
  • ...ไปเป็น tuple ของ list...
  • ...ไปเป็น tuple ของ PyTorch tensor
Deep Reinforcement Learning ด้วย Python

นำ Experience Replay มาใช้ใน DQN

  1. ก่อน training loop: replay_buffer = ReplayBuffer(10000)

  2. ใน training loop หลังเลือก 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)
  • กำหนดค่าเริ่มต้น replay buffer
  • เพิ่ม transition ล่าสุดเข้า buffer

หากความยาว buffer $\geq$ batch_size:

  • สุ่ม batch จาก buffer แล้วคำนวณ loss
  • การคำนวณ loss ไม่เปลี่ยนแปลงในเชิงแนวคิด
  • Mean Squared Bellman Error บน batch จาก replay memory
    • การเรียนรู้มีเสถียรภาพและประสิทธิภาพมากขึ้น
Deep Reinforcement Learning ด้วย Python

มาฝึกกันเถอะ!

Deep Reinforcement Learning ด้วย Python

Preparing Video For Download...