Q-ținte fixe
Te pregătești să antrenezi Lunar Lander cu Q-ținte fixe. Ca pas pregătitor, trebuie să instanțiezi atât rețeaua online (care alege acțiunea), cât și rețeaua țintă (folosită pentru calculul țintei TD).
Trebuie, de asemenea, să implementezi o funcție update_target_network pe care o vei putea apela la fiecare pas de antrenament. Rețeaua țintă nu este actualizată prin coborâre pe gradient; în schimb, update_target_network îi ajustează ușor ponderile în direcția rețelei Q, asigurând astfel că rămâne stabilă în timp.
Reține că, doar pentru acest exercițiu, folosești o rețea foarte mică, astfel încât să poți printa și inspecta cu ușurință dicționarul său de stări. Are un singur strat ascuns de dimensiune doi; spațiul acțiunilor și spațiul stărilor sunt, de asemenea, de dimensiune 2.
Funcția print_state_dict() este disponibilă în mediul tău pentru a afișa dicționarul de stări.
Acest exercițiu face parte din cursul
Deep Reinforcement Learning în Python
Instrucțiuni pentru exercițiu
- Obține
.state_dict()atât pentru rețeaua țintă, cât și pentru cea online. - Actualizează dicționarul de stări al rețelei țintă calculând media ponderată dintre parametrii rețelei online și cei ai rețelei țintă, folosind
tauca pondere pentru rețeaua online. - Încarcă dicționarul de stări actualizat înapoi în rețeaua țintă.
Exercițiu interactiv practic
Încearcă acest exercițiu completând acest cod de exemplu.
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))