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() và .__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
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êmself.epsilonđể bao quát các trường hợp biên. - Trong
.increase_beta(), tăng beta lênself.beta_increment; đảm bảobetakhô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)