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

Bộ đệm prioritized experience replay

Bạn sẽ giới thiệu lớp PrioritizedExperienceReplay, một cấu trúc dữ liệu mà bạn sẽ dùng sau đó để triển khai DQN với Prioritized Experience Replay.

PrioritizedExperienceReplay là một cải tiến so với lớp ExperienceReplay mà bạn đã dùng đến giờ để huấn luyện các agent DQN. Một bộ đệm prioritized experience replay đảm bảo rằng các transition được lấy mẫu từ đó có giá trị học tập cao hơn cho agent so với lấy mẫu đồng đều.

Hiện tại, hãy triển khai các phương thức .__init__(), .push(), .update_priorities(), .increase_beta().__len__(). Phương thức cuối cùng, .sample(), sẽ là trọng tâm của Bài tập tiếp theo.

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

  • Trong .push(), khởi tạo mức độ ưu tiên của transition bằng giá trị ưu tiên lớn nhất trong bộ đệm (hoặc 1 nếu bộ đệm trống).
  • Trong .update_priorities(), đặt mức độ ưu tiên bằng giá trị tuyệt đối của TD error tương ứng; cộng thêm self.epsilon để bao quát các trường hợp biên.
  • Trong .increase_beta(), tăng beta lên self.beta_increment; đảm bảo beta không vượt quá 1.

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.

class PrioritizedReplayBuffer:
    def __init__(
        self, capacity, alpha=0.6, beta=0.4, beta_increment=0.001, epsilon=0.01
    ):
        self.memory = deque(maxlen=capacity)
        self.alpha, self.beta, self.beta_increment, self.epsilon = (alpha, beta, beta_increment, epsilon)
        self.priorities = deque(maxlen=capacity)

    def push(self, state, action, reward, next_state, done):
        experience_tuple = (state, action, reward, next_state, done)
        # Initialize the transition's priority
        max_priority = ____
        self.memory.append(experience_tuple)
        self.priorities.append(max_priority)
    
    def update_priorities(self, indices, td_errors):
        for idx, td_error in zip(indices, td_errors.tolist()):
            # Update the transition's priority
            self.priorities[idx] = ____

    def increase_beta(self):
        # Increase beta if less than 1
        self.beta = ____

    def __len__(self):
        return len(self.memory)
      
buffer = PrioritizedReplayBuffer(capacity=3)
buffer.push(state=[1,3], action=2, reward=1, next_state=[2,4], done=False)
print("Transition in memory buffer:", buffer.memory)
print("Priority buffer:", buffer.priorities)
Chỉnh sửa và Chạy Mã