บัฟเฟอร์ Prioritized Experience Replay
ในแบบฝึกหัดนี้ จะได้สร้างคลาส PrioritizedExperienceReplay ซึ่งเป็นโครงสร้างข้อมูลที่จะนำไปใช้ในการพัฒนา DQN แบบ Prioritized Experience Replay ในภายหลัง
PrioritizedExperienceReplay คือการปรับปรุงจากคลาส ExperienceReplay ที่ใช้ฝึก DQN agent มาโดยตลอด บัฟเฟอร์แบบ Prioritized Experience Replay จะคัดเลือก transition ที่มีประโยชน์ต่อการเรียนรู้ของ agent มากกว่าการสุ่มแบบ uniform
ตอนนี้ให้ implement เมธอด .__init__(), .push(), .update_priorities(), .increase_beta() และ .__len__() ส่วนเมธอดสุดท้ายคือ .sample() จะนำไปฝึกในแบบฝึกหัดถัดไป
แบบฝึกหัดนี้เป็นส่วนหนึ่งของหลักสูตร
Deep Reinforcement Learning ด้วย Python
คำแนะนำการฝึกหัด
- ใน
.push()ให้กำหนดค่าลำดับความสำคัญเริ่มต้นของ transition เป็นค่าสูงสุดในบัฟเฟอร์ (หรือ 1 หากบัฟเฟอร์ยังว่างอยู่) - ใน
.update_priorities()ให้กำหนดค่าลำดับความสำคัญเป็นค่าสัมบูรณ์ของ TD error ที่สอดคล้องกัน แล้วบวกself.epsilonเพื่อรองรับกรณีพิเศษ - ใน
.increase_beta()ให้เพิ่มค่า beta ด้วยself.beta_incrementโดยตรวจสอบว่าbetaต้องไม่เกิน 1
แบบฝึกหัดเชิงโต้ตอบแบบลงมือทำ
ลองทำแบบฝึกหัดนี้โดยเติมโค้ดตัวอย่างนี้ให้สมบูรณ์
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)