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

DQN ร่วมกับ Prioritized Experience Replay

ในแบบฝึกหัดนี้ จะนำ Prioritized Experience Replay (PER) มาใช้เพื่อปรับปรุงอัลกอริทึม DQN โดย PER มีเป้าหมายเพื่อเพิ่มประสิทธิภาพการเลือกชุด transition ที่ใช้อัปเดตโครงข่ายในแต่ละขั้นตอน

สำหรับอ้างอิง เมธอดที่ประกาศไว้ใน PrioritizedReplayBuffer มีดังนี้:

  • push() (สำหรับเพิ่ม transition เข้าบัฟเฟอร์)
  • sample() (สำหรับสุ่มชุด transition จากบัฟเฟอร์)
  • increase_beta() (สำหรับเพิ่มค่า importance sampling)
  • update_priorities() (สำหรับอัปเดตค่า priority ที่สุ่มได้)

ฟังก์ชัน describe_episode() ถูกนำมาใช้อีกครั้งเพื่ออธิบายแต่ละ episode

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

Deep Reinforcement Learning ด้วย Python

ดูคอร์ส

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

  • สร้าง Prioritized Experience Replay buffer ที่มีความจุ 10000 transitions
  • เพิ่มอิทธิพลของ importance sampling ตามเวลาโดยอัปเดตพารามิเตอร์ beta
  • อัปเดตค่า priority ของ experience ที่สุ่มได้ตาม TD error ล่าสุดของแต่ละรายการ

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

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

# Instantiate a Prioritized Replay Buffer with capacity 10000
replay_buffer = ____(____)

for episode in range(5):
    state, info = env.reset()
    done = False   
    step = 0
    episode_reward = 0    
    # Increase the replay buffer's beta parameter
    replay_buffer.____
    while not done:
        step += 1
        total_steps += 1
        q_values = online_network(state)
        action = select_action(q_values, total_steps, start=.9, end=.05, decay=1000)
        next_state, reward, terminated, truncated, _ = env.step(action)
        done = terminated or truncated
        replay_buffer.push(state, action, reward, next_state, done)        
        if len(replay_buffer) >= batch_size:
            states, actions, rewards, next_states, dones, indices, weights = replay_buffer.sample(64)
            q_values = online_network(states).gather(1, actions).squeeze(1)
            with torch.no_grad():
                next_q_values = target_network(next_states).amax(1)
                target_q_values = rewards + gamma * next_q_values * (1-dones)            
            td_errors = target_q_values - q_values
            # Update the replay buffer priorities for that batch
            replay_buffer.____(____, ____)
            loss = torch.sum(weights * (q_values - target_q_values) ** 2)
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            update_target_network(target_network, online_network, tau=.005)
        state = next_state
        episode_reward += reward    
    describe_episode(episode, reward, episode_reward, step)
แก้ไขและรันโค้ด