La fonction objectif substitut tronquée
Implémentez la fonction calculate_loss() pour PPO. Il s'agit de coder l'innovation clé de PPO — la fonction de perte substitutive tronquée. Elle sert à limiter la mise à jour de la politique pour éviter qu'elle ne s'éloigne trop de la politique précédente à chaque étape.
La formule de l'objectif substitut tronqué est
Votre environnement a l'hyperparamètre de troncature epsilon réglé à 0.2.
Cette activité fait partie du cours
Deep Reinforcement Learning en Python
Instructions de l’exercice
- Obtenez les rapports de probabilités entre
\pi_\thetaet\pi_{\theta_{old}}(versions non tronquée et tronquée). - Calculez les objectifs substituts (versions non tronquée et tronquée).
- Calculez l'objectif substitut tronqué de PPO.
- Calculez la perte de l'acteur.
Exercice interactif pratique
Essayez cet exercice en complétant ce code d’exemple.
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)