Replay-buffer för erfarenheter
Nu ska du skapa den datastruktur som krävs för Experience Replay, vilket gör att agenten kan lära sig betydligt effektivare.
Replay-bufferten ska stödja två operationer:
- Lagra erfarenheter i sitt minne för framtida sampling.
- "Spela upp" ett slumpmässigt sampel av tidigare erfarenheter från sitt minne.
Eftersom data som samplas från replay-bufferten används som indata till ett neuralt nätverk bör bufferten returnera torch-tensorer för enkel hantering.
Modulerna torch och random samt klassen deque har importerats till din övningsmiljö.
Den här övningen är en del av kursen
Djup förstärkningsinlärning i Python
Övningsinstruktioner
- Komplettera metoden
push()iReplayBuffergenom att lägga tillexperience_tuplei buffertens minne. - I metoden
sample(), dra ett slumpmässigt sampel av storlekenbatch_sizefrånself.memory. - Fortfarande i
sample(): samplet hämtas initialt som en lista av tupler – se till att det omvandlas till en tupel av listor. - Omvandla
actions_tensortill formen(batch_size, 1)i stället för(batch_size).
Interaktiv övning med praktiskt arbete
Testa den här övningen genom att slutföra den här exempelkoden.
class ReplayBuffer:
def __init__(self, capacity):
self.memory = deque([], maxlen=capacity)
def push(self, state, action, reward, next_state, done):
experience_tuple = (state, action, reward, next_state, done)
# Append experience_tuple to the memory buffer
self.memory.____
def __len__(self):
return len(self.memory)
def sample(self, batch_size):
# Draw a random sample of size batch_size
batch = ____(____, ____)
# Transform batch into a tuple of lists
states, actions, rewards, next_states, dones = ____
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)
# Ensure actions_tensor has shape (batch_size, 1)
actions_tensor = torch.tensor(actions, dtype=torch.long).____
return states_tensor, actions_tensor, rewards_tensor, next_states_tensor, dones_tensor