ÎncepețiÎncepe gratuit

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

Vezi cursul

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 tau ca 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))
Editează și rulează codul