开始使用免费开始使用

最简版 DQN 损失函数

现在 select_action() 函数已就绪,距离训练智能体只差最后一步:实现 calculate_loss()

calculate_loss() 会返回该回合任一步的网络损失。

参考公式如下:

本练习已加载以下示例数据:

state = torch.rand(8)
next_state = torch.rand(8)
action = select_action(q_network, state)
reward = 1
gamma = .99
done = False

本练习是课程的一部分

Python 中的深度强化学习

查看课程

练习说明

  • 获取当前状态的 Q 值。
  • 获取下一状态的 Q 值。
  • 计算目标 Q 值(TD-target)。
  • 计算损失函数,即平方贝尔曼误差。

交互式实操练习

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

def calculate_loss(q_network, state, action, next_state, reward, done):
    q_values = q_network(state)
    print(f'Q-values: {q_values}')
    # Obtain the current state Q-value
    current_state_q_value = q_values[____]
    print(f'Current state Q-value: {current_state_q_value:.2f}')
    # Obtain the next state Q-value
    next_state_q_value = q_network(next_state).____    
    print(f'Next state Q-value: {next_state_q_value:.2f}')
    # Calculate the target Q-value
    target_q_value = ____ + gamma * ____ * (1-done)
    print(f'Target Q-value: {target_q_value:.2f}')
    # Obtain the loss
    loss = nn.MSELoss()(____, ____)
    print(f'Loss: {loss:.2f}')
    return loss

calculate_loss(q_network, state, action, next_state, reward, done)
编辑并运行代码