Modifying the loss in ppo in stable-baselines3

Viewed 237

I'm trying to implement an addition to the loss function of the ppo algorithm in stable-baselines3. For this I collected additional observations for the states s(t-10) and s(t+1) which I can access in the train-function of the PPO class in ppo.py as part of the rollout_buffer.

I'm using a 3-layer-mlp as my network architecture and need the outputs of the second layer for the triplet (s(t-α), s(t), s(t+1)) to use them to calculate L = max(d(s(t+1) , s(t)) − d(s(t+1) , s(t−α)) + γ, 0), where d is the L2-distance.

Finally I want to add this term to the old loss, so loss = loss + 0.3 * L

This is my implementation starting with the original loss in line 242:

                loss = policy_loss + self.ent_coef * entropy_loss + self.vf_coef * value_loss

                ###############################

                net1 = nn.Sequential(*list(self.policy.mlp_extractor.policy_net.children())[:-1])
                L_losses = []
                a = 0
                obs = rollout_data.observations
                obs_alpha = rollout_data.observations_alpha 
                obs_plusone = rollout_data.observations_plusone
                inds = rollout_data.inds

                for i in inds:
                    if i > alpha: # only use observations for which L can be calculated
                        fs_t = net1(obs[a])
                        fs_talpha = net1(obs_alpha[a])
                        fs_tone = net1(obs_plusone[a])
                        L = max(
                            th.norm(th.subtract(fs_tone, fs_t)) - th.norm(th.subtract(fs_tone, fs_talpha)) + 1.0, 0.0)
                        L_losses.append(L)
                    else:
                        L_losses.append(0)
                    a += 1
                L_loss = th.mean(th.FloatTensor(L_losses))
                loss += 0.3 * L_loss

So with net1 I tried to get a clone of the original network with the outputs from the second layer. I am unsure if this is the right way to do this.

I do have some questions about my approach as the resulting performance is slightly worse compared to without the added term although it should be slightly better:

  1. Is my way of getting the outputs of the second layer of the mlp network working?
  2. When loss.backward() is called can the gradient be calculated correctly (with the new term included)?
0 Answers
Related