Fixed Q-targets
आप Fixed Q-targets के साथ अपने Lunar Lander को ट्रेन करने की तैयारी कर रहे हैं. पूर्वापेक्षा के तौर पर, आपको ऑनलाइन नेटवर्क (जो एक्शन चुनता है) और टार्गेट नेटवर्क (जिसका उपयोग TD-target कैलकुलेशन के लिए होता है) दोनों को इंस्टैनशिएट करना है.
आपको एक update_target_network फंक्शन भी इम्प्लीमेंट करना है, जिसका उपयोग आप हर ट्रेनिंग स्टेप पर कर सकेंगे. टार्गेट नेटवर्क को gradient descent से अपडेट नहीं किया जाता; बल्कि, update_target_network उसके वेट्स को थोड़ी-थोड़ी मात्रा में Q-network की ओर धकेलता है, ताकि समय के साथ यह पर्याप्त रूप से स्थिर बना रहे.
ध्यान दें कि, केवल इस अभ्यास के लिए, आप एक बहुत छोटा नेटवर्क उपयोग कर रहे हैं ताकि हम आसानी से उसका state dictionary प्रिंट और निरीक्षण कर सकें. इसमें सिर्फ एक hidden लेयर है जिसका साइज दो है; इसका action space और state space भी dimension 2 का है.
आपके वातावरण में print_state_dict() फंक्शन उपलब्ध है जो state dict प्रिंट करता है.
यह अभ्यास पाठ्यक्रम का हिस्सा है
Python में Deep Reinforcement Learning
अभ्यास निर्देश
- टार्गेट और ऑनलाइन, दोनों नेटवर्क्स के लिए
.state_dict()प्राप्त करें. tauको ऑनलाइन नेटवर्क के वेट के रूप में लेते हुए, ऑनलाइन और टार्गेट नेटवर्क्स के पैरामीटर्स का weighted average लेकर टार्गेट नेटवर्क के 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))