Fixed Q-targets
ขณะนี้คุณกำลังเตรียมฝึก Lunar Lander ด้วย fixed Q-targets โดยก่อนอื่น จำเป็นต้องสร้าง online network (ซึ่งใช้เลือก action) และ target network (ซึ่งใช้คำนวณ TD-target) ขึ้นมาก่อน
นอกจากนี้ ยังต้องสร้างฟังก์ชัน update_target_network ที่จะเรียกใช้ในแต่ละขั้นตอนการฝึก target network จะไม่ถูกอัปเดตด้วย gradient descent แต่ฟังก์ชัน update_target_network จะค่อย ๆ ปรับค่าน้ำหนักให้เข้าใกล้ Q-network ทีละน้อย เพื่อให้ target network มีความเสถียรตลอดการฝึก
ในแบบฝึกหัดนี้จะใช้โครงข่ายขนาดเล็กเพื่อให้ตรวจสอบ state dictionary ได้ง่าย โดยมี hidden layer เพียงหนึ่งชั้นที่มีขนาด 2 และทั้ง action space กับ state space มีมิติเท่ากับ 2
ฟังก์ชัน print_state_dict() พร้อมใช้งานในสภาพแวดล้อมของคุณสำหรับแสดงผล state dict
แบบฝึกหัดนี้เป็นส่วนหนึ่งของหลักสูตร
Deep Reinforcement Learning ด้วย Python
คำแนะนำการฝึกหัด
- ดึง
.state_dict()ของทั้ง target network และ online network - อัปเดต state dict ของ target network โดยคำนวณค่าเฉลี่ยถ่วงน้ำหนักระหว่างพารามิเตอร์ของ online network กับ target network โดยใช้
tauเป็นค่าน้ำหนักของ online network - โหลด state dict ที่อัปเดตแล้วกลับเข้าสู่ target network
แบบฝึกหัดเชิงโต้ตอบแบบลงมือทำ
ลองทำแบบฝึกหัดนี้โดยเติมโค้ดตัวอย่างนี้ให้สมบูรณ์
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))