Actualizări batch în policy gradient

Deep Reinforcement Learning în Python

Timothée Carayol

Principal Machine Learning Engineer, Komment

Actualizări gradient: pas cu pas vs. batch

O casetă mare reprezentând un episod.

Deep Reinforcement Learning în Python

Actualizări gradient: pas cu pas vs. batch

În caseta mare apare o casetă mai mică reprezentând pasul 1. În interior, o casetă cu textul 'selectare acțiune'.

Deep Reinforcement Learning în Python

Actualizări gradient: pas cu pas vs. batch

În caseta pasului 1 apare o altă casetă mică cu textul 'iterare mediu'

Deep Reinforcement Learning în Python

Actualizări gradient: pas cu pas vs. batch

Sub caseta pasului 1, o altă casetă cu etichetele 'calculare pierdere' și 'gradient descent'

Deep Reinforcement Learning în Python

Actualizări gradient: pas cu pas vs. batch

Un al doilea set identic de casete apare pentru pasul 2, cu același conținut

Deep Reinforcement Learning în Python

Actualizări gradient: pas cu pas vs. batch

Apar și etapele 3 și 4.

Deep Reinforcement Learning în Python

Actualizări batch A2C / PPO

O casetă mare de episod; ocupând jumătate din suprafața sa, o casetă etichetată 'rollout 1'; în interior, două casete goale etichetate 'pasul 1' și 'pasul 2'

Deep Reinforcement Learning în Python

Actualizări batch A2C / PPO

În caseta pasului 1 apar etichetele 'selectare acțiune' și 'iterare mediu'.

Deep Reinforcement Learning în Python

Actualizări batch A2C / PPO

La fel pentru pasul 2.

Deep Reinforcement Learning în Python

Actualizări batch A2C / PPO

Sub casetele pasului 1 și pasului 2 apare o singură etichetă 'calculare pierdere' și o singură etichetă 'gradient descent'.

Deep Reinforcement Learning în Python

Actualizări batch A2C / PPO

Cealaltă jumătate a episodului este acum ocupată de un al doilea rollout identic cu două etape, etichetat 'rollout 2'.

Deep Reinforcement Learning în Python

Bucla de antrenament A2C cu actualizări batch

 

# Set rollout length
rollout_length = 10

# Initiate loss batches
actor_losses = torch.tensor([]) critic_losses = torch.tensor([])
  • Inițializare batch-uri de pierdere
  • Iterare prin episoade și pași ca de obicei

 

for episode in range(10):
  state, info = env.reset()
  done = False
  while not done:
    action, action_log_prob = select_action(actor, 
                                            state)                
    next_state, reward, terminated, truncated, _ = (
                                   env.step(action))
    done = terminated or truncated    
    actor_loss, critic_loss = calculate_losses(
        critic, action_log_prob, 
        reward, state, next_state, done)
    ...
Deep Reinforcement Learning în Python

Bucla de antrenament A2C cu actualizări batch

  ...
  actor_losses = torch.cat((actor_losses, actor_loss))
  critic_losses = torch.cat((critic_losses, critic_loss))

# If rollout is full, update the networks if len(actor_losses) >= rollout_length:
actor_loss_batch = actor_losses.mean() critic_loss_batch = critic_losses.mean()
actor_optimizer.zero_grad() actor_loss_batch.backward() actor_optimizer.step() critic_optimizer.zero_grad() critic_loss_batch.backward() critic_optimizer.step()
actor_losses = torch.tensor([]) critic_losses = torch.tensor([])
state = next_state

 

  • Adăugare pierdere pas la batch-urile de pierdere
  • Când rollout-ul este complet:
    • Calculare medie batch cu .mean()
    • Efectuare gradient descent
    • Reinițializare batch-uri de pierdere
Deep Reinforcement Learning în Python

A2C / PPO cu mai mulți agenți

 

Două benzi orizontale reprezentând agentul 1 și agentul 2. Fiecare agent parcurge respectiv 4 și 3 episoade de lungimi variabile. În interiorul fiecărui episod sunt vizibile casete de pași. Sub cele două benzi sunt vizibile trei casete de rollout, fiecare acoperind un interval de 8 pași. În fiecare casetă de rollout sunt vizibile etichetele 'calculare pierdere' și 'gradient descent'. În partea de sus, o legendă indică: „lungime rollout: 8 pași; număr de agenți: 2"

Deep Reinforcement Learning în Python

Rollout-uri și minibatch-uri

Două benzi de agenți identice cu slide-ul anterior. Dedesubt, 3 casete de rollout sunt din nou vizibile, dar conținutul lor s-a schimbat. Fiecare are în partea de sus o casetă lungă etichetată 'amestecare'. Sub aceasta, sunt împărțite longitudinal în 4 casete etichetate 'minibatch'; în fiecare minibatch se află o casetă 'calculare pierdere' și una 'gradient descent'. Legenda indică: 'Lungime rollout: 8 pași; dimensiune minibatch: 4 (2x2); număr de agenți: 2'

Deep Reinforcement Learning în Python

PPO cu mai multe epoci

Un desen similar cu cel anterior, cu excepția casetelor de rollout, care sunt acum împărțite vertical în 4 zone: prima este o etichetă 'amestecare'; a doua este o casetă mare etichetată 'epoca 1' conținând 4 minibatch-uri dispuse longitudinal; a treia este o etichetă 'reamestecare'; ultima este o casetă mare etichetată 'epoca 2', conținând tot 4 minibatch-uri. Legenda indică: 'Lungime rollout: 8 pași; dimensiune minibatch: 4 (2x2); număr de agenți: 2; număr de epoci: 2'.

Deep Reinforcement Learning în Python

Să exersăm!

Deep Reinforcement Learning în Python

Preparing Video For Download...