Začněte nyníZačněte zdarma

Fixní Q-targety

Chystáš se trénovat Lunar Lander s fixními Q-targety. Nejprve je potřeba vytvořit online síť (která vybírá akci) i cílovou síť (používanou pro výpočet TD-targetu).

Musíš také implementovat funkci update_target_network, kterou budeš volat při každém tréninkovém kroku. Cílová síť se neaktualizuje gradientním sestupem – místo toho ji update_target_network postupně přibližuje váhám Q-sítě o malý krok, takže zůstává v čase stabilní.

Pozor: jen pro toto cvičení pracuješ s velmi malou sítí, aby bylo snadné vypsat a prozkoumat její state dict. Má jedinou skrytou vrstvu o velikosti 2; action space i state space mají také dimenzi 2.

Funkce print_state_dict() je v prostředí k dispozici pro výpis state dictu.

Toto cvičení je součástí kurzu

Deep Reinforcement Learning v Pythonu

Zobrazit kurz

Pokyny k cvičení

  • Získej .state_dict() jak cílové, tak online sítě.
  • Aktualizuj state dict cílové sítě jako vážený průměr parametrů online sítě a cílové sítě – jako váhu online sítě použij tau.
  • Načti aktualizovaný state dict zpět do cílové sítě.

Interaktivní cvičení na vyzkoušení si v praxi

Vyzkoušejte si toto cvičení dokončením tohoto ukázkového kódu.

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))
Upravit a spustit kód