शुरू करेंमुफ़्त में शुरू करें

Barebone 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 में Deep Reinforcement Learning

पाठ्यक्रम देखें

अभ्यास निर्देश

  • वर्तमान स्टेट का Q-वैल्यू प्राप्त करें.
  • अगले स्टेट का Q-वैल्यू प्राप्त करें.
  • टारगेट Q-वैल्यू (या TD-target) की गणना करें.
  • लॉस फंक्शन की गणना करें, यानी squared Bellman Error.

इंटरैक्टिव व्यावहारिक अभ्यास

इस अभ्यास को इस नमूना कोड को पूरा करके आज़माएँ।

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)
कोड संपादित करें और चलाएँ