始める無料で始める

固定Qターゲット

固定Qターゲットを用いて Lunar Lander を学習させる準備をします。前提として、行動を選択するオンラインネットワークと、TDターゲットの計算に用いるターゲットネットワークの両方をインスタンス化する必要があります。

また、各学習ステップで使用できる update_target_network 関数を実装します。ターゲットネットワークは勾配降下で更新しません。代わりに、update_target_network はその重みを少しだけ Q-network に近づけることで、時間とともに十分安定した状態を保つようにします。

この演習に限り、state dictionary を簡単に表示・確認できるように、非常に小さなネットワークを使います。隠れ層はサイズ2が1層のみで、アクション空間と状態空間の次元も2です。

環境には、state dict を出力するための print_state_dict() 関数が用意されています。

この演習はコースの一部です

Pythonで学ぶDeep Reinforcement Learning

コースを見る

演習の手順

  • ターゲットネットワークとオンラインネットワークそれぞれの .state_dict() を取得します。
  • オンラインネットワークの重みを tau として、オンラインとターゲットの各パラメータの加重平均を取り、ターゲットネットワーク用の state dict を更新します。
  • 更新した state dict をターゲットネットワークに読み込みます。

実践的なインタラクティブ演習

このサンプルコードを完成させて、この演習に挑戦してみましょう。

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))
コードを編集して実行