НачатьНачать бесплатно

Расчёт функций потерь Actor-Critic

Перед тем как перейти к обучению агента с помощью A2C, напишите функцию calculate_losses(), которая возвращает значения потерь для обеих сетей.

Для справки — ниже приведены выражения для функций потерь актора и критика соответственно:

Это упражнение является частью курса

Глубокое обучение с подкреплением на Python

Посмотреть курс

Инструкции к упражнению

  • Вычислите TD-цель.
  • Вычислите функцию потерь для сети-актора.
  • Вычислите функцию потерь для сети-критика.

Интерактивное практическое упражнение

Попробуйте выполнить это упражнение, дополнив этот пример кода.

def calculate_losses(critic_network, action_log_prob, 
                     reward, state, next_state, done):
    value = critic_network(state)
    next_value = critic_network(next_state)
    # Calculate the TD target
    td_target = (____ + gamma * ____ * (1-done))
    td_error = td_target - value
    # Calculate the actor loss
    actor_loss = -____ * ____.detach()
    # Calculate the critic loss
    critic_loss = ____
    return actor_loss, critic_loss
  
actor_loss, critic_loss = calculate_losses(
        critic_network, action_log_prob, 
        reward, state, next_state, done
)
print(round(actor_loss.item(), 2), round(critic_loss.item(), 2))
Редактировать и запускать код