Stałe cele Q (Fixed Q-targets)
Przygotowujesz się do trenowania lądownika księżycowego z użyciem stałych celów Q. Na początku musisz zainicjalizować dwie sieci: sieć online (która wybiera akcje) oraz sieć docelową (używaną do obliczania celu TD).
Musisz również zaimplementować funkcję update_target_network, którą będziesz wywoływać na każdym kroku trenowania. Sieć docelowa nie jest aktualizowana przez gradient descent – zamiast tego update_target_network stopniowo przesuwa jej wagi w kierunku sieci Q o niewielką wartość, dzięki czemu pozostaje stabilna przez cały czas trenowania.
Zwróć uwagę, że tylko w tym ćwiczeniu używana jest bardzo mała sieć, co pozwala łatwo wydrukować i przeanalizować jej słownik stanu. Ma ona tylko jedną ukrytą warstwę o rozmiarze 2; przestrzeń akcji i przestrzeń stanów mają również wymiar 2.
Funkcja print_state_dict() jest dostępna w twoim środowisku i służy do wydrukowania słownika stanu.
To ćwiczenie jest częścią kursu
Głębokie uczenie ze wzmocnieniem w Pythonie
Instrukcje do ćwiczenia
- Pobierz
.state_dict()zarówno dla sieci docelowej, jak i sieci online. - Zaktualizuj słownik stanu sieci docelowej, obliczając ważoną średnią parametrów sieci online i sieci docelowej, używając
taujako wagi dla sieci online. - Wczytaj zaktualizowany słownik stanu z powrotem do sieci docelowej.
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
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))