НачатьНачать бесплатно

Буфер приоритетного воспроизведения опыта

Вы реализуете класс PrioritizedExperienceReplay — структуру данных, которую в дальнейшем используете для обучения DQN с приоритетным воспроизведением опыта.

PrioritizedExperienceReplay — это улучшенная версия класса ExperienceReplay, который вы применяли ранее для обучения агентов DQN. Буфер приоритетного воспроизведения опыта гарантирует, что переходы, извлекаемые из него, оказываются более ценными для обучения агента по сравнению с равномерной выборкой.

Пока что реализуйте методы .__init__(), .push(), .update_priorities(), .increase_beta() и .__len__(). Последний метод, .sample(), станет темой следующего упражнения.

Это упражнение является частью курса

Глубокое обучение с подкреплением на Python

Посмотреть курс

Инструкции к упражнению

  • В методе .push() инициализируйте приоритет перехода максимальным приоритетом в буфере (или значением 1, если буфер пуст).
  • В методе .update_priorities() установите приоритет равным абсолютному значению соответствующей TD-ошибки; прибавьте self.epsilon, чтобы обработать граничные случаи.
  • В методе .increase_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)
Редактировать и запускать код