ÎncepețiÎncepe gratuit

A2C cu actualizări pe batch

Pe parcursul acestui curs, ai folosit variante ale aceluiași ciclu de antrenament DRL de bază. În practică, există mai multe moduri în care această structură poate fi extinsă, de exemplu pentru a include actualizări pe batch.

Vei relua ciclul de antrenament A2C pe mediul Lunar Lander, dar în loc să actualizezi rețelele la fiecare pas, vei aștepta 10 pași înainte de a rula pasul de coborâre pe gradient. Prin medierea pierderilor pe 10 pași, vei obține actualizări ușor mai stabile.

Acest exercițiu face parte din cursul

Deep Reinforcement Learning în Python

Vezi cursul

Instrucțiuni pentru exercițiu

  • Adaugă pierderile din fiecare pas la tensorii de pierderi pentru batch-ul curent.
  • Calculează pierderile pe batch.
  • Reinițializează tensorii de pierderi.

Exercițiu interactiv practic

Încearcă acest exercițiu completând acest cod de exemplu.

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)
Editează și rulează codul