Буфер воспроизведения опыта
Сейчас вы создадите структуру данных для поддержки воспроизведения опыта (Experience Replay), которая позволит агенту обучаться значительно эффективнее.
Буфер воспроизведения должен поддерживать две операции:
- Сохранение опыта в памяти для последующей выборки.
- «Воспроизведение» случайно выбранного пакета прошлых примеров из памяти.
Поскольку данные из буфера воспроизведения будут подаваться на вход нейронной сети, буфер должен возвращать тензоры torch для удобства работы.
Модули torch и random, а также класс deque уже импортированы в среду выполнения упражнения.
Это упражнение является частью курса
Глубокое обучение с подкреплением на Python
Инструкции к упражнению
- Завершите реализацию метода
push()классаReplayBuffer, добавивexperience_tupleв память буфера. - В методе
sample()извлеките случайную выборку размеромbatch_sizeизself.memory. - Там же в
sample()выборка изначально получается в виде списка кортежей; преобразуйте её в кортеж списков. - Преобразуйте
actions_tensorтак, чтобы его форма была(batch_size, 1)вместо(batch_size).
Интерактивное практическое упражнение
Попробуйте выполнить это упражнение, дополнив этот пример кода.
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