Kom igångKom igång gratis

Sampling från PER-bufferten

Innan du kan använda klassen Prioritized Experience Buffer för att träna din agent behöver du fortfarande implementera metoden .sample(). Metoden tar storleken på det urval du vill dra som argument och returnerar de samplade övergångarna som tensors, tillsammans med deras index i minnesbufferten och deras importansvikter.

En buffert med kapacitet 10 har laddats in i din miljö så att du kan sampla från den.

Den här övningen är en del av kursen

Djup förstärkningsinlärning i Python

Visa kurs

Övningsinstruktioner

  • Beräkna samplingssannolikheten för varje övergång.
  • Dra de index som motsvarar övergångarna i urvalet; np.random.choice(a, s, p=p) tar ett urval av storleken s med återläggning från arrayen a, baserat på sannolikhetsarrayen p.
  • Beräkna importansvikten för varje övergång.

Interaktiv övning med praktiskt arbete

Testa den här övningen genom att slutföra den här exempelkoden.

def sample(self, batch_size):
    priorities = np.array(self.priorities)
    # Calculate the sampling probabilities
    probabilities = ____ / np.sum(____)
    # Draw the indices for the sample
    indices = np.random.choice(____)
    # Calculate the importance weights
    weights = (1 / (len(self.memory) * ____)) ** ____
    weights /= np.max(weights)
    states, actions, rewards, next_states, dones = zip(*[self.memory[idx] for idx in indices])
    weights = [weights[idx] for idx in indices]
    states_tensor = torch.tensor(states, dtype=torch.float32)
    rewards_tensor = torch.tensor(rewards, dtype=torch.float32)
    next_states_tensor = torch.tensor(next_states, dtype=torch.float32)
    dones_tensor = torch.tensor(dones, dtype=torch.float32)
    weights_tensor = torch.tensor(weights, dtype=torch.float32)
    actions_tensor = torch.tensor(actions, dtype=torch.long).unsqueeze(1)
    return (states_tensor, actions_tensor, rewards_tensor, next_states_tensor,
            dones_tensor, indices, weights_tensor)

PrioritizedReplayBuffer.sample = sample
print("Sampled transitions:\n", buffer.sample(3))
Redigera och kör kod