Полный алгоритм DQN

Глубокое обучение с подкреплением на Python

Timothée Carayol

Principal Machine Learning Engineer, Komment

Алгоритм DQN

 

 

  • Мы изучили DQN с буфером воспроизведения опыта
  • Близко к оригинальному DQN (2015)
  • Нам не хватает двух компонентов:
    • Epsilon-жадность -> больше исследования
    • Фиксированные Q-цели -> более стабильное обучение

Путешественница с греческой буквой эпсилон рядом

Заглавная буква Q, замороженная в блоке льда

Глубокое обучение с подкреплением на Python

Epsilon-жадность в алгоритме DQN

  • Реализуйте убывающую Epsilon-жадность в select_action()
def select_action(q_values, step, start, end, decay):

# Calculate the threshold value for this step epsilon = ( end + (start-end) * math.exp(-step / decay))
# Draw a random number between 0 and 1 sample = random.random()
if sample < epsilon: # Return a random action index return random.choice(range(len(q_values)))
# Return the action index with highest Q-value return torch.argmax(q_values).item()
  • $\varepsilon = end + (start-end) \cdot e^{-\frac{step}{decay}}$
  • Случайное действие выбирается с вероятностью $\varepsilon$
  • Действие с максимальным значением — с вероятностью $1 - \varepsilon$

График убывания эпсилон для различных значений параметра decay.

Глубокое обучение с подкреплением на Python

Фиксированные Q-цели

 

  • В ошибке Беллмана:
    • Q-сеть используется и для Q-значений, и для TD-цели
    • Нестабильность из-за смещения цели

 

  • Целевая сеть стабилизирует цель

 

Ошибка Беллмана: (r_t+1 + gamma max(Q(s_t+1, a))) - Q(s_t, a_t)

(r_t+1 + gamma max(Q_target(s_t+1, a))) - Q_online(s_t, a_t)

Глубокое обучение с подкреплением на Python

Реализация фиксированных Q-целей

online_network = QNetwork(state_size, action_size)
target_network = QNetwork(state_size, action_size)

target_network.load_state_dict( online_network.state_dict())
def update_target_network( target_network, online_network, tau):
target_net_state_dict = target_network.state_dict() online_net_state_dict = online_network.state_dict() for key in online_net_state_dict:
target_net_state_dict[key] = ( online_net_state_dict[key] * tau + target_net_state_dict[key] * (1 - tau))
target_network.load_state_dict( target_net_state_dict)
return None
  • Изначально онлайн-сеть = целевая сеть
  • Словарь состояния сети содержит все веса: Представление словаря состояния сети с записями fc1.weight, fc1.bias и fc2.weight; значение каждой записи — тензор.
  • На каждом шаге веса целевой сети постепенно приближаются к онлайн-сети
Глубокое обучение с подкреплением на Python

Расчёт потерь с фиксированными Q-целями

# In the inner loop, after action selection
if len(replay_buffer) >= batch_size:
  states, actions, rewards, next_states, dones = 
      replay_buffer.sample(64)

q_values = (online_network(states) .gather(1, actions).squeeze(1))
with torch.no_grad():
next_q_values = ( target_network(next_states).amax(1)) target_q_values = ( rewards + gamma * next_q_values * (1 - dones))
loss = torch.nn.MSELoss()(target_q_values, q_values) optimizer.zero_grad() loss.backward() optimizer.step()
update_target_network( target_network, online_network, tau)

 

  • Q-значения вычисляются через online_network
  • Целевые Q-значения — через target_network
  • torch.no_grad() отключает отслеживание градиентов для целевых Q-значений
  • Для расчёта потерь по-прежнему используется среднеквадратичная ошибка Беллмана
  • update_target_network() постепенно обновляет target_network
Глубокое обучение с подкреплением на Python

Давайте потренируемся!

Глубокое обучение с подкреплением на Python

Preparing Video For Download...