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
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))