Zacznij terazZacznij za darmo

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

Zobacz kurs

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 tau jako 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))
Edytuj i uruchom kod