НачатьНачать бесплатно

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

Вы готовитесь к обучению агента Lunar Lander с использованием фиксированных Q-целей. В качестве предварительного шага необходимо создать две сети: основную (которая выбирает действие) и целевую (используемую для вычисления TD-цели).

Также нужно реализовать функцию update_target_network, которую вы будете вызывать на каждом шаге обучения. Целевая сеть не обновляется методом градиентного спуска — вместо этого update_target_network постепенно сдвигает её веса в сторону Q-сети на небольшую величину, обеспечивая её стабильность на протяжении всего обучения.

Обратите внимание: только в этом упражнении используется очень небольшая сеть, чтобы было удобно печатать и изучать её словарь состояния. В ней всего один скрытый слой размером два; пространство действий и пространство состояний также имеют размерность 2.

Функция print_state_dict() доступна в вашем окружении для вывода словаря состояния.

Это упражнение является частью курса

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

Посмотреть курс

Инструкции к упражнению

  • Получите .state_dict() для целевой и основной сетей.
  • Обновите словарь состояния целевой сети, вычислив взвешенное среднее между параметрами основной и целевой сетей: используйте tau как вес для основной сети.
  • Загрузите обновлённый словарь состояния обратно в целевую сеть.

Интерактивное практическое упражнение

Попробуйте выполнить это упражнение, дополнив этот пример кода.

def update_target_network(target_network, online_network, tau):
    # Obtain the state dicts for both networks
    target_net_state_dict = ____
    online_net_state_dict = ____
    for key in online_net_state_dict:
        # Calculate the updated state dict for the target network
        target_net_state_dict[key] = (online_net_state_dict[____] * ____ + target_net_state_dict[____] * ____)
        # Load the updated state dict into the target network
        target_network.____
    return None
  
print("online network weights:", print_state_dict(online_network))
print("target network weights (pre-update):", print_state_dict(target_network))
update_target_network(target_network, online_network, .001)
print("target network weights (post-update):", print_state_dict(target_network))
Редактировать и запускать код