Bắt đầu ngayBắt đầu miễn phí

A2C với cập nhật theo batch

Cho đến lúc này trong khóa học, bạn đã dùng nhiều biến thể xoay quanh cùng một vòng lặp huấn luyện DRL cốt lõi. Thực tế có nhiều cách mở rộng cấu trúc này, ví dụ để hỗ trợ cập nhật theo batch.

Giờ bạn sẽ xem lại vòng lặp huấn luyện A2C trên môi trường Lunar Lander, nhưng thay vì cập nhật mạng ở mỗi bước, bạn sẽ đợi cho đến khi trôi qua 10 bước rồi mới thực hiện bước hạ gradient. Bằng cách lấy trung bình loss trong 10 bước, bạn sẽ có các cập nhật ổn định hơn một chút.

Bài tập này là một phần của khóa học

Deep Reinforcement Learning bằng Python

Xem khóa học

Hướng dẫn bài tập

  • Nối các loss từ mỗi bước vào các tensor loss cho batch hiện tại.
  • Tính batch loss.
  • Khởi tạo lại các tensor loss.

Bài tập tương tác thực hành trực tiếp

Hãy thử làm bài tập này bằng cách hoàn thành đoạn mã mẫu này.

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)
Chỉnh sửa và Chạy Mã