Zacznij terazZacznij za darmo

Bufor z priorytetyzowanym powtarzaniem doświadczeń

Zaimplementujesz klasę PrioritizedExperienceReplay – strukturę danych, którą później wykorzystasz do zbudowania DQN z priorytetyzowanym powtarzaniem doświadczeń.

PrioritizedExperienceReplay to udoskonalona wersja klasy ExperienceReplay, której używałeś do tej pory podczas trenowania agentów DQN. Bufor z priorytetyzowanym powtarzaniem doświadczeń gwarantuje, że próbkowane z niego przejścia są bardziej wartościowe dla uczenia się agenta niż te wybierane losowo z jednakowym prawdopodobieństwem.

Na razie zaimplementuj metody .__init__(), .push(), .update_priorities(), .increase_beta() oraz .__len__(). Ostatnia metoda, .sample(), będzie tematem kolejnego ćwiczenia.

To ćwiczenie jest częścią kursu

Głębokie uczenie ze wzmocnieniem w Pythonie

Zobacz kurs

Instrukcje do ćwiczenia

  • W metodzie .push() zainicjalizuj priorytet przejścia jako maksymalny priorytet w buforze (lub 1, jeśli bufor jest pusty).
  • W metodzie .update_priorities() ustaw priorytet jako wartość bezwzględną odpowiadającego błędu TD; dodaj self.epsilon, aby obsłużyć przypadki brzegowe.
  • W metodzie .increase_beta() zwiększ beta o self.beta_increment; zadbaj o to, żeby beta nigdy nie przekroczyło 1.

Interaktywne ćwiczenie praktyczne

Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.

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)
Edytuj i uruchom kod