시작하기무료로 시작하기

기본 DQN 손실 함수

이제 select_action() 함수를 준비했으니, 에이전트를 학습시키기까지 마지막 한 단계만 남았어요. 바로 calculate_loss()를 구현하는 것입니다.

calculate_loss()는 에피소드의 임의의 단계에 대해 네트워크 손실을 반환합니다.

참고로, 손실은 다음과 같이 주어집니다:

다음 예시 데이터가 연습 문제에 로드되어 있습니다:

state = torch.rand(8)
next_state = torch.rand(8)
action = select_action(q_network, state)
reward = 1
gamma = .99
done = False

이 연습은 강의의 일부입니다

Python으로 배우는 Deep Reinforcement Learning

강의 보기

연습 안내

  • 현재 상태의 Q-값을 구하세요.
  • 다음 상태의 Q-값을 구하세요.
  • 목표 Q-값(TD-target)을 계산하세요.
  • 손실 함수, 즉 제곱 Bellman 오차를 계산하세요.

실습형 인터랙티브 연습

이 예제를 이 샘플 코드를 완성하여 풀어보세요.

def calculate_loss(q_network, state, action, next_state, reward, done):
    q_values = q_network(state)
    print(f'Q-values: {q_values}')
    # Obtain the current state Q-value
    current_state_q_value = q_values[____]
    print(f'Current state Q-value: {current_state_q_value:.2f}')
    # Obtain the next state Q-value
    next_state_q_value = q_network(next_state).____    
    print(f'Next state Q-value: {next_state_q_value:.2f}')
    # Calculate the target Q-value
    target_q_value = ____ + gamma * ____ * (1-done)
    print(f'Target Q-value: {target_q_value:.2f}')
    # Obtain the loss
    loss = nn.MSELoss()(____, ____)
    print(f'Loss: {loss:.2f}')
    return loss

calculate_loss(q_network, state, action, next_state, reward, done)
코드 편집 및 실행