Kom igångKom igång gratis

A2C med batchuppdateringar

I den här kursen har du hittills arbetat med varianter av samma grundläggande DRL-träningsloop. I praktiken finns det flera sätt att bygga ut den strukturen – till exempel för att hantera batchuppdateringar.

Du kommer nu att återbesöka A2C-träningsloopen i Lunar Lander-miljön, men i stället för att uppdatera nätverken vid varje steg väntar du tills 10 steg har genomförts innan du kör gradientnedstigningssteget. Genom att beräkna medelvärdet av förlusterna över 10 steg får du något mer stabila uppdateringar.

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

Djup förstärkningsinlärning i Python

Visa kurs

Övningsinstruktioner

  • Lägg till förlusterna från varje steg i förlustensorerna för den aktuella batchen.
  • Beräkna batchförlusterna.
  • Återinitialisera förlustensorerna.

Interaktiv övning med praktiskt arbete

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

actor_losses = torch.tensor([])
critic_losses = torch.tensor([])
for episode in range(10):
    state, info = env.reset()
    done = False
    episode_reward = 0
    step = 0
    while not done:
        step += 1
        action, action_log_prob = select_action(actor, state)                
        next_state, reward, terminated, truncated, _ = env.step(action)
        done = terminated or truncated
        episode_reward += reward
        actor_loss, critic_loss = calculate_losses(
            critic, action_log_prob, 
            reward, state, next_state, done)
        # Append to the loss tensors
        actor_losses = torch.cat((____, ____))
        critic_losses = torch.cat((____, ____))
        if len(actor_losses) >= 10:
            # Calculate the batch losses
            actor_loss_batch = actor_losses.____
            critic_loss_batch = critic_losses.____
            actor_optimizer.zero_grad(); actor_loss_batch.backward(); actor_optimizer.step()
            critic_optimizer.zero_grad(); critic_loss_batch.backward(); critic_optimizer.step()
            # Reinitialize the loss tensors
            actor_losses = ____
            critic_losses = ____
        state = next_state
    describe_episode(episode, reward, episode_reward, step)
Redigera och kör kod