시작하기무료로 시작하기

클리핑된 대리 목적 함수

PPO용 calculate_loss() 함수를 구현하세요. 이는 PPO의 핵심 혁신인 클리핑된 대리 손실 함수를 코드로 작성하는 작업이에요. 이 손실은 각 단계에서 정책이 이전 정책에서 너무 멀리 벗어나지 않도록 정책 업데이트를 제한해 줍니다.

클리핑된 대리 목적 함수의 수식은 다음과 같습니다.

이 환경에서는 클리핑 하이퍼파라미터 epsilon이 0.2로 설정되어 있어요.

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

Python으로 배우는 Deep Reinforcement Learning

강의 보기

연습 안내

  • \pi_\theta\pi_{\theta_{old}} 사이의 확률 비율을 구하세요(클리핑 전/후 버전 모두).
  • 대리 목적 함수를 계산하세요(클리핑 전/후 버전 모두).
  • PPO 클리핑된 대리 목적 함수를 계산하세요.
  • actor 손실을 계산하세요.

실습형 인터랙티브 연습

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

def calculate_losses(critic_network, action_log_prob, action_log_prob_old,
                     reward, state, next_state, done):
    value = critic_network(state)
    next_value = critic_network(next_state)
    td_target = (reward + gamma * next_value * (1-done))
    td_error = td_target - value
    # Obtain the probability ratios
    ____, ____ = calculate_ratios(action_log_prob, action_log_prob_old, epsilon=.2)
    # Calculate the surrogate objectives
    surr1 = ratio * ____.____()
    surr2 = clipped_ratio * ____.____()    
    # Calculate the clipped surrogate objective
    objective = torch.min(____, ____)
    # Calculate the actor loss
    actor_loss = ____
    critic_loss = td_error ** 2
    return actor_loss, critic_loss
  
actor_loss, critic_loss = calculate_losses(critic_network, action_log_prob, action_log_prob_old,
                                           reward, state, next_state, done)
print(actor_loss, critic_loss)
코드 편집 및 실행