경험 재생을 활용한 DQN

Python으로 배우는 Deep Reinforcement Learning

Timothée Carayol

Principal Machine Learning Engineer, Komment

경험 재생 소개

 

  • 기본 DQN 에이전트는 최신 경험만 학습
    • 연속 업데이트 간 상관이 큼
    • 금방 잊어버림
  • 해결: 경험 재생
    • 경험을 버퍼에 저장
    • 매 스텝마다 과거에서 무작위 배치로 학습

 

울타리 미로의 항공 사진

Python으로 배우는 Deep Reinforcement Learning

양끝 큐(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])
  • 용량 초과 시 가장 오래된 항목이 제거됨

용량 7의 deque. 1~4 요소 포함.

용량 7의 deque. 1~4 요소, 오른쪽에 5~8 추가 예정.

용량 7의 deque. 1~7 요소 포함. 오른쪽에 8 대기.

용량 7의 deque. 2~8 요소 포함. 왼쪽에 1이 밀려남.

Python으로 배우는 Deep Reinforcement Learning

리플레이 버퍼 구현

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)
...
  • 리플레이 메모리: 제한 용량의 deque
  • .push():
    • 전이를 튜플로 저장
    • 버퍼에 추가
    • 용량 도달 시 가장 오래된 전이 제거
Python으로 배우는 Deep Reinforcement Learning

리플레이 버퍼 구현

...
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: 전이 튜플 리스트에서...
  • ...리스트의 튜플로...
  • ...PyTorch 텐서의 튜플로 변환
Python으로 배우는 Deep Reinforcement Learning

DQN에 경험 재생 통합

  1. 학습 루프 전: replay_buffer = ReplayBuffer(10000)

  2. 학습 루프에서, 행동 선택 후:

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)
  • 리플레이 버퍼 초기화
  • 최신 전이를 버퍼에 추가

버퍼 길이 $\geq$ batch_size이면:

  • 버퍼에서 무작위 배치를 뽑아 손실 계산 진행
  • 손실 계산 개념은 동일
  • 리플레이 배치에 대한 평균제곱 벨만 오차
    • 학습이 더 안정적이고 효율적임
Python으로 배우는 Deep Reinforcement Learning

Vamos praticar!

Python으로 배우는 Deep Reinforcement Learning

Preparing Video For Download...