Фиксированные 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))