Buffer de memorie pentru Experience Replay
Vei crea acum structura de date necesară pentru Experience Replay, care va permite agentului tău să învețe mult mai eficient.
Acest buffer de memorie trebuie să suporte două operații:
- Stocarea experiențelor în memorie pentru eșantionare ulterioară.
- „Redarea" unui lot de experiențe trecute, eșantionat aleatoriu din memorie.
Deoarece datele extrase din buffer vor fi folosite ca intrare într-o rețea neuronală, buffer-ul ar trebui să returneze tensori torch pentru comoditate.
Modulele torch și random, precum și clasa deque, au fost deja importate în mediul de lucru.
Acest exercițiu face parte din cursul
Deep Reinforcement Learning în Python
Instrucțiuni pentru exercițiu
- Completează metoda
push()a claseiReplayBufferprin adăugarea luiexperience_tupleîn memoria buffer-ului. - În metoda
sample(), extrage un eșantion aleatoriu de dimensiunebatch_sizedinself.memory. - Tot în
sample(), eșantionul este extras inițial ca o listă de tupluri; asigură-te că este transformat într-un tuplu de liste. - Transformă
actions_tensorastfel încât să aibă forma(batch_size, 1)în loc de(batch_size).
Exercițiu interactiv practic
Încearcă acest exercițiu completând acest cod de exemplu.
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