Bufor powtórek doświadczeń
Czas stworzyć strukturę danych obsługującą mechanizm powtórek doświadczeń (Experience Replay), dzięki której agent będzie uczył się znacznie efektywniej.
Ten bufor powtórek powinien obsługiwać dwie operacje:
- Zapisywanie doświadczeń w pamięci z myślą o przyszłym losowaniu.
- „Odtwarzanie" losowo wybranej paczki (batcha) przeszłych doświadczeń z pamięci.
Ponieważ dane pobrane z bufora będą podawane do sieci neuronowej, bufor powinien zwracać tensory torch dla wygody.
Moduły torch i random oraz klasa deque zostały już zaimportowane do środowiska ćwiczenia.
To ćwiczenie jest częścią kursu
Głębokie uczenie ze wzmocnieniem w Pythonie
Instrukcje do ćwiczenia
- Uzupełnij metodę
push()klasyReplayBuffer, dodającexperience_tupledo pamięci bufora. - W metodzie
sample()wylosuj próbkę o rozmiarzebatch_sizezself.memory. - Nadal w metodzie
sample()– próbka jest początkowo pobierana jako lista krotek; upewnij się, że zostanie przekształcona w krotkę list. - Przekształć
actions_tensortak, aby miał kształt(batch_size, 1)zamiast(batch_size).
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
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