Kom igångKom igång gratis

Förlustfunktion för ett grundläggande DQN

Nu när funktionen select_action() är klar är du bara ett steg ifrån att kunna träna din agent: du ska nu implementera calculate_loss().

Funktionen calculate_loss() returnerar nätverkets förlust för ett givet steg i episoden.

För referens ges förlusten av:

Följande exempeldata har laddats in i övningen:

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

Den här övningen är en del av kursen

Djup förstärkningsinlärning i Python

Visa kurs

Övningsinstruktioner

  • Hämta det aktuella tillståndets Q-värde.
  • Hämta nästa tillstånds Q-värde.
  • Beräkna mål-Q-värdet, det vill säga TD-målet.
  • Beräkna förlustfunktionen, det vill säga det kvadratiska Bellman-felet.

Interaktiv övning med praktiskt arbete

Testa den här övningen genom att slutföra den här exempelkoden.

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)
Redigera och kör kod