Функция потерь базового 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-цель).
- Вычислите функцию потерь, то есть квадратичную ошибку Беллмана.
Интерактивное практическое упражнение
Попробуйте выполнить это упражнение, дополнив этот пример кода.
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)