Отсечённая суррогатная целевая функция
Реализуйте функцию calculate_loss() для PPO. Для этого необходимо запрограммировать ключевое нововведение PPO — отсечённую суррогатную функцию потерь. Она ограничивает обновление политики, не позволяя ей слишком сильно отклоняться от предыдущей политики на каждом шаге.
Формула отсечённой суррогатной цели:
В вашей среде гиперпараметр отсечения epsilon установлен равным 0.2.
Это упражнение является частью курса
Глубокое обучение с подкреплением на Python
Инструкции к упражнению
- Вычислите отношения вероятностей между
\pi_\thetaи\pi_{\theta_{old}}(нескорректированную и скорректированную версии). - Рассчитайте суррогатные цели (нескорректированную и скорректированную версии).
- Рассчитайте отсечённую суррогатную цель 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)