Hàm mất mát DQN bản tối giản
Với hàm select_action() đã sẵn sàng, bạn chỉ còn một bước nữa là có thể huấn luyện agent: bây giờ bạn sẽ hiện thực calculate_loss().
Hàm calculate_loss() trả về giá trị mất mát của mạng tại một bước bất kỳ trong episode.
Tham khảo, công thức mất mát được cho bởi:
Ví dụ dữ liệu sau đã được nạp sẵn trong bài tập:
state = torch.rand(8)
next_state = torch.rand(8)
action = select_action(q_network, state)
reward = 1
gamma = .99
done = False
Bài tập này là một phần của khóa học
Deep Reinforcement Learning bằng Python
Hướng dẫn bài tập
- Lấy Q-value của trạng thái hiện tại.
- Lấy Q-value của trạng thái kế tiếp.
- Tính Q-value mục tiêu (TD-target).
- Tính hàm mất mát, tức là Sai số Bellman bình phương.
Bài tập tương tác thực hành trực tiếp
Hãy thử làm bài tập này bằng cách hoàn thành đoạn mã mẫu này.
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)