เริ่มต้นใช้งานเริ่มต้นใช้งานได้ฟรี

A2C กับการอัปเดตแบบ batch

ตลอดคอร์สนี้ เราได้ใช้รูปแบบการเทรน DRL หลักที่คล้ายกันในหลายรูปแบบ ในทางปฏิบัติ โครงสร้างนี้สามารถขยายได้หลายวิธี เช่น การรองรับการอัปเดตแบบ batch

คราวนี้จะกลับมาดู training loop ของ A2C บนสภาพแวดล้อม Lunar Lander อีกครั้ง แต่แทนที่จะอัปเดตโครงข่ายประสาทเทียมทุก step จะรอให้ครบ 10 step ก่อนจึงค่อยรัน gradient descent การเฉลี่ย loss ข้าม 10 step จะช่วยให้การอัปเดตมีความเสถียรมากขึ้น

แบบฝึกหัดนี้เป็นส่วนหนึ่งของหลักสูตร

Deep Reinforcement Learning ด้วย Python

ดูคอร์ส

คำแนะนำการฝึกหัด

  • Append ค่า loss จากแต่ละ step ลงใน tensor ที่เก็บ loss ของ batch ปัจจุบัน
  • คำนวณ batch loss
  • กำหนดค่าเริ่มต้นใหม่ให้กับ tensor ที่เก็บ loss

แบบฝึกหัดเชิงโต้ตอบแบบลงมือทำ

ลองทำแบบฝึกหัดนี้โดยเติมโค้ดตัวอย่างนี้ให้สมบูรณ์

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)
แก้ไขและรันโค้ด