Advantage Actor Critic

Deep Reinforcement Learning v Pythonu

Timothée Carayol

Principal Machine Learning Engineer, Komment

Proč Actor Critic?

 

  • Omezení REINFORCE:

    • Vysoký rozptyl
    • Nízká vzorková efektivita
  • Metody Actor Critic přidávají critic síť, která umožňuje učení s Temporal Difference

Velký obdélník označený „agent"; uvnitř dva menší obdélníky označené „actor" a „critic".

Deep Reinforcement Learning v Pythonu

Intuice za metodami Actor Critic

Studenti diskutují u stolu, kolem jsou rozloženy knihy a tužky.

 

  • Actor síť:

    • Rozhoduje
    • Neumí svá rozhodnutí hodnotit
  • Critic síť:

    • V každém kroku poskytuje actoru zpětnou vazbu
Deep Reinforcement Learning v Pythonu

Critic síť

 

  • Critic aproximuje funkci hodnoty stavu

Znázornění critic sítě se stavem jako vstupem a hodnotovou funkcí jako výstupem; síť má pouze jeden výstupní uzel.

  • Hodnotí akci $a_t$ na základě advantage nebo TD chyby

 

class Critic(nn.Module):
    def __init__(self, state_size):
        super(Critic, self).__init__()
        self.fc1 = nn.Linear(state_size, 64)
        self.fc2 = nn.Linear(64, 1)

def forward(self, state): x = torch.relu(self.fc1(torch.tensor(state))) value = self.fc2(x) return value
critic_network = Critic(8)
Deep Reinforcement Learning v Pythonu

Dynamika Actor Critic

 

  • V každém kroku:
    • Actor volí akci (stejně jako v REINFORCE)

Nahoře: velký obdélník označený „agent"; uvnitř dva menší obdélníky označené „actor" a „critic". Dole: samostatný obdélník označený „environment".

Deep Reinforcement Learning v Pythonu

Dynamika Actor Critic

 

  • V každém kroku:
    • Actor volí akci (stejně jako v REINFORCE)
    • Critic pozoruje odměnu a stav

Červená šipka označená „action" vede od actora k prostředí.

Deep Reinforcement Learning v Pythonu

Dynamika Actor Critic

 

  • V každém kroku:
    • Actor volí akci (stejně jako v REINFORCE)
    • Critic pozoruje odměnu a stav
    • Critic vyhodnocuje TD Error
    • Actor i Critic aktualizují váhy pomocí TD Error

Dvě červené šipky označené „State" a „Reward" vedou z prostředí do Criticu.

Deep Reinforcement Learning v Pythonu

Dynamika Actor Critic

 

  • V každém kroku:
    • Actor volí akci (stejně jako v REINFORCE)
    • Critic pozoruje odměnu a stav
    • Critic vyhodnocuje TD Error
    • Actor i Critic aktualizují váhy pomocí TD Error
    • Aktualizovaný Actor pozoruje nový stav

Šipka označená „TD error" vede od Criticu k Actoru.

Deep Reinforcement Learning v Pythonu

Dynamika Actor Critic

 

  • V každém kroku:
    • Actor volí akci (stejně jako v REINFORCE)
    • Critic pozoruje odměnu a stav
    • Critic vyhodnocuje TD Error
    • Actor i Critic aktualizují váhy pomocí TD Error
    • Aktualizovaný Actor pozoruje nový stav
  • ... opakovat

Šipka State nyní směřuje také k Actoru.

Deep Reinforcement Learning v Pythonu

Ztráty A2C

 

Critic

Ztrátová funkce criticu. Používá se kvadratická TD chyba: Lc(theta c) = ((r_t + gamma * V theta c (s t + 1)) - V theta c) na druhou

  • Ztráta criticu: kvadratická TD chyba

 

Actor

Ztrátová funkce actoru. V každém kroku t lze použít: L(theta) se rovná minus logaritmus pravděpodobnosti akce krát TD chyba nebo advantage.

  • TD chyba zachycuje hodnocení criticu
  • Zvyšuje pravděpodobnost akcí s kladnou TD chybou
Deep Reinforcement Learning v Pythonu

Výpočet ztrát

 

def calculate_losses(critic_network, action_log_prob, 
                     reward, state, next_state, done):

# Critic provides the state value estimates value = critic_network(state)
next_value = critic_network(next_state)
td_target = (reward + gamma * next_value * (1-done))
td_error = td_target - value
# Apply formulas for actor and critic losses actor_loss = -action_log_prob * td_error.detach()
critic_loss = td_error ** 2
return actor_loss, critic_loss

 

 

  • Výpočet TD chyby
  • Výpočet ztráty actoru
    • .detach() zastaví propagaci gradientu do vah criticu
  • Výpočet ztráty criticu
Deep Reinforcement Learning v Pythonu

Trénovací smyčka Actor Critic

for episode in range(10):
  state, info = env.reset()
  done = False
  while not done:

# Select action action, action_log_prob = select_action(actor, state)
next_state, reward, terminated, truncated, _ = env.step(action) done = terminated or truncated
# Calculate losses actor_loss, critic_loss = calculate_losses(critic, action_log_prob, reward, state, next_state, done)
# Update actor actor_optimizer.zero_grad(); actor_loss.backward(); actor_optimizer.step()
# Update critic critic_optimizer.zero_grad(); critic_loss.backward(); critic_optimizer.step()
state = next_state
Deep Reinforcement Learning v Pythonu

Pojďme cvičit!

Deep Reinforcement Learning v Pythonu

Preparing Video For Download...