완전한 DQN 알고리즘

Python으로 배우는 Deep Reinforcement Learning

Timothée Carayol

Principal Machine Learning Engineer, Komment

DQN 알고리즘

 

 

  • 리플레이 버퍼가 있는 DQN을 학습했습니다
  • 최초 발표(2015)의 DQN과 유사합니다
  • 아직 두 구성 요소가 남았습니다:
    • 엡실론-탐욕(epsilon-greediness) → 탐색 증가
    • 고정 Q-타깃 → 학습 안정화

모험가가 서 있고, 그 옆에 그리스 문자 엡실론이 있음

얼음 블록에 얼어 있는 대문자 Q

Python으로 배우는 Deep Reinforcement Learning

DQN의 엡실론-탐욕

  • select_action()에 감쇠 엡실론-탐욕 구현
def select_action(q_values, step, start, end, decay):

# 이 스텝의 임계값 계산 epsilon = ( end + (start-end) * math.exp(-step / decay))
# 0과 1 사이 난수 추출 sample = random.random()
if sample < epsilon: # 무작위 액션 인덱스 반환 return random.choice(range(len(q_values)))
# Q-value가 최대인 액션 인덱스 반환 return torch.argmax(q_values).item()
  • $\varepsilon = end + (start-end) \cdot e^{-\frac{step}{decay}}$
  • 확률 $\varepsilon$로 무작위 행동 선택
  • 확률 $1 - \varepsilon$로 최대 값 행동 선택

감쇠 파라미터 값별 엡실론 감쇠 스케줄 그래프

Python으로 배우는 Deep Reinforcement Learning

고정 Q-타깃

 

  • 벨만 오류에서:
    • Q-Value와 TD-타깃 모두에 Q-Network 사용
    • 타깃 변화로 인한 불안정성

 

  • 타깃 네트워크로 타깃을 안정화

 

벨만 오류: (r_t+1 + gamma max(Q(s_t+1, a))) - Q(s_t, a_t)

(r_t+1 + gamma max(Q_target(s_t+1, a))) - Q_online(s_t, a_t)

Python으로 배우는 Deep Reinforcement Learning

고정 Q-타깃 구현

online_network = QNetwork(state_size, action_size)
target_network = QNetwork(state_size, action_size)

target_network.load_state_dict( online_network.state_dict())
def update_target_network( target_network, online_network, tau):
target_net_state_dict = target_network.state_dict() online_net_state_dict = online_network.state_dict() for key in online_net_state_dict:
target_net_state_dict[key] = ( online_net_state_dict[key] * tau + target_net_state_dict[key] * (1 - tau))
target_network.load_state_dict( target_net_state_dict)
return None
  • 처음에는 Online Network = Target Network
  • 네트워크의 state dict에는 모든 가중치가 포함됨: 네트워크 state dictionary의 예: fc1.weight, fc1.bias, fc2.weight 항목과 각 항목의 텐서 값.
  • 매 스텝마다 Target Network의 각 가중치를 Online Network에 조금씩 근접시킴
Python으로 배우는 Deep Reinforcement Learning

고정 Q-타깃으로 손실 계산

# 내부 루프, 액션 선택 후
if len(replay_buffer) >= batch_size:
  states, actions, rewards, next_states, dones = 
      replay_buffer.sample(64)

q_values = (online_network(states) .gather(1, actions).squeeze(1))
with torch.no_grad():
next_q_values = ( target_network(next_states).amax(1)) target_q_values = ( rewards + gamma * next_q_values * (1 - dones))
loss = torch.nn.MSELoss()(target_q_values, q_values) optimizer.zero_grad() loss.backward() optimizer.step()
update_target_network( target_network, online_network, tau)

 

  • Q-value는 online_network 사용
  • Target Q-value는 target_network 사용
  • 타깃 Q-value에는 torch.no_grad()로 그래디언트 비활성화
  • 손실은 여전히 평균제곱 벨만 오류 사용
  • update_target_network()target_network를 천천히 갱신
Python으로 배우는 Deep Reinforcement Learning

Lass uns üben!

Python으로 배우는 Deep Reinforcement Learning

Preparing Video For Download...