ПочатиПочніть безкоштовно

Clipped surrogate objective — цільова функція з обрізанням

Реалізуйте функцію calculate_loss() для PPO. Для цього потрібно закодувати ключову інновацію PPO — функцію втрат на основі clipped surrogate. Вона обмежує оновлення політики, щоб на кожному кроці політика не відхилялася надто далеко від попередньої.

Формула для clipped surrogate objective така:

У вашому середовищі гіперпараметр обрізання epsilon встановлено на 0.2.

Ця вправа є частиною курсу

Глибоке навчання з підкріпленням у Python

Переглянути курс

Інструкції до вправи

  • Отримайте відношення імовірностей між \pi_\theta та \pi_{\theta_{old}} (некліпована та кліпована версії).
  • Обчисліть замінникові (surrogate) цілі (некліповану та кліповану версії).
  • Обчисліть clipped surrogate objective для PPO.
  • Обчисліть втрати актора.

Інтерактивна практична вправа

Спробуйте виконати цю вправу, доповнивши цей зразок коду.

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)
Редагувати та запускати код