DQN с приоритизированным воспроизведением опыта
В этом упражнении вы внедрите приоритизированное воспроизведение опыта (Prioritized Experience Replay, PER), чтобы улучшить алгоритм DQN. PER позволяет оптимизировать формирование батча переходов, используемого для обновления сети на каждом шаге.
Для справки: методы, объявленные для PrioritizedReplayBuffer:
push()— добавляет переходы в буфер;sample()— выбирает батч переходов из буфера;increase_beta()— увеличивает вес взвешенной выборки по важности;update_priorities()— обновляет приоритеты выбранных переходов.
Функция describe_episode() используется, как и прежде, для описания каждого эпизода.
Это упражнение является частью курса
Глубокое обучение с подкреплением на Python
Инструкции к упражнению
- Создайте буфер приоритизированного воспроизведения опыта ёмкостью 10000 переходов.
- Постепенно увеличивайте влияние взвешенной выборки по важности, обновляя параметр
beta. - Обновите приоритеты выбранных переходов на основе их последних TD-ошибок.
Интерактивное практическое упражнение
Попробуйте выполнить это упражнение, дополнив этот пример кода.
# 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)