Den klippta surrogat-målfunktionen
Implementera funktionen calculate_loss() för PPO. Det innebär att du kodar PPO:s viktigaste innovation – den klippta surrogat-förlustfunktionen. Den hjälper till att begränsa policyuppdateringen så att den inte avviker för mycket från den tidigare policyn vid varje steg.
Formeln för det klippta surrogatmålet är
Din miljö har klippningshyperparametern epsilon satt till 0,2.
Den här övningen är en del av kursen
Djup förstärkningsinlärning i Python
Övningsinstruktioner
- Beräkna sannolikhetskvoterna mellan
\pi_\thetaoch\pi_{\theta_{old}}(oklippt och klippt version). - Beräkna surrogatmålen (oklippt och klippt version).
- Beräkna det klippta PPO-surrogatmålet.
- Beräkna aktörförlusten.
Interaktiv övning med praktiskt arbete
Testa den här övningen genom att slutföra den här exempelkoden.
def calculate_losses(critic_network, action_log_prob, action_log_prob_old,
reward, state, next_state, done):
value = critic_network(state)
next_value = critic_network(next_state)
td_target = (reward + gamma * next_value * (1-done))
td_error = td_target - value
# Obtain the probability ratios
____, ____ = calculate_ratios(action_log_prob, action_log_prob_old, epsilon=.2)
# Calculate the surrogate objectives
surr1 = ratio * ____.____()
surr2 = clipped_ratio * ____.____()
# Calculate the clipped surrogate objective
objective = torch.min(____, ____)
# Calculate the actor loss
actor_loss = ____
critic_loss = td_error ** 2
return actor_loss, critic_loss
actor_loss, critic_loss = calculate_losses(critic_network, action_log_prob, action_log_prob_old,
reward, state, next_state, done)
print(actor_loss, critic_loss)