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