开始使用免费开始使用

Actor-Critic 的损失计算

在使用 A2C 训练智能体之前的最后一步,请编写一个 calculate_losses() 函数,返回两个网络各自的损失。

供参考,以下分别是 Actor 与 Critic 的损失函数表达式:

本练习是课程的一部分

Python 中的深度强化学习

查看课程

练习说明

  • 计算 TD 目标。
  • 计算 Actor 网络的损失。
  • 计算 Critic 网络的损失。

交互式实操练习

通过完成这段示例代码来试试这个练习。

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))
编辑并运行代码