Bắt đầu ngayBắt đầu miễn phí

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

Xem khóa học

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)
Chỉnh sửa và Chạy Mã