Prioritized experience replay के साथ DQN
इस अभ्यास में, आप DQN एल्गोरिदम में सुधार करने के लिए Prioritized Experience Replay (PER) जोड़ेंगे. PER का लक्ष्य हर स्टेप पर नेटवर्क अपडेट करने के लिए चुने गए ट्रांज़िशन्स के बैच को बेहतर बनाना है.
संदर्भ के लिए, PrioritizedReplayBuffer के लिए आपने जिन मेथड नामों की घोषणा की है, वे हैं:
push()(ट्रांज़िशन्स को बफ़र में भेजने के लिए)sample()(बफ़र से ट्रांज़िशन्स का एक बैच सैंपल करने के लिए)increase_beta()(इम्पॉर्टेंस सैंपलिंग का प्रभाव बढ़ाने के लिए)update_priorities()(सैंपल की गई प्रायोरिटीज़ अपडेट करने के लिए)
हर एपिसोड का विवरण देने के लिए describe_episode() फंक्शन फिर से उपयोग किया गया है.
यह अभ्यास पाठ्यक्रम का हिस्सा है
Python में Deep Reinforcement Learning
अभ्यास निर्देश
- 10000 ट्रांज़िशन्स की क्षमता वाला एक Prioritized Experience Replay बफ़र इंस्टैंशिएट करें.
- समय के साथ
betaपैरामीटर अपडेट करके इम्पॉर्टेंस सैंपलिंग के प्रभाव को बढ़ाएँ. - सैंपल की गई एक्सपीरियंसेज़ की नवीनतम 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)