@@ -328,15 +328,15 @@ def _forward_learn(self, data: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
328328 # discrete part (discrete policy loss and entropy loss)
329329 ppo_discrete_batch = ppo_policy_data (
330330 output ['logit' ]['action_type' ], batch ['logit' ]['action_type' ], batch ['action' ]['action_type' ],
331- adv , batch ['weight' ]
331+ adv , batch ['weight' ], logit_pretrained
332332 )
333333 ppo_discrete_loss , ppo_discrete_info = ppo_policy_error (
334334 ppo_discrete_batch , self ._clip_ratio , kl_type = self ._kl_type
335335 )
336336 # continuous part (continuous policy loss and entropy loss, value loss)
337337 ppo_continuous_batch = ppo_data (
338338 output ['logit' ]['action_args' ], batch ['logit' ]['action_args' ], batch ['action' ]['action_args' ],
339- output ['value' ], batch ['value' ], adv , batch ['return' ], batch ['weight' ]
339+ output ['value' ], batch ['value' ], adv , batch ['return' ], batch ['weight' ], None
340340 )
341341 ppo_continuous_loss , ppo_continuous_info = ppo_error_continuous (
342342 ppo_continuous_batch , self ._clip_ratio , kl_type = self ._kl_type
@@ -794,7 +794,7 @@ def _forward_learn(self, data: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
794794 output = self ._learn_model .forward (batch ['obs' ])
795795
796796 ppo_batch = ppo_policy_data (
797- output ['logit' ], batch ['logit' ], batch ['action' ], batch ['return' ], batch ['weight' ]
797+ output ['logit' ], batch ['logit' ], batch ['action' ], batch ['return' ], batch ['weight' ], None
798798 )
799799 if self ._action_space == 'continuous' :
800800 ppo_loss , ppo_info = ppo_policy_error_continuous (ppo_batch , self ._clip_ratio )
@@ -1235,26 +1235,26 @@ def _forward_learn(self, data: List[Dict[str, Any]]) -> Dict[str, Any]:
12351235 if self ._action_space == 'continuous' :
12361236 ppodata = ppo_data (
12371237 output ['logit' ], data ['logit' ], data ['action' ], output ['value' ], data ['value' ], adv , data ['return' ],
1238- data ['weight' ]
1238+ data ['weight' ], None
12391239 )
12401240 ppo_loss , ppo_info = ppo_error_continuous (ppodata , self ._clip_ratio )
12411241 elif self ._action_space == 'discrete' :
12421242 ppodata = ppo_data (
12431243 output ['logit' ], data ['logit' ], data ['action' ], output ['value' ], data ['value' ], adv , data ['return' ],
1244- data ['weight' ]
1244+ data ['weight' ], None
12451245 )
12461246 ppo_loss , ppo_info = ppo_error (ppodata , self ._clip_ratio )
12471247 elif self ._action_space == 'hybrid' :
12481248 # discrete part (discrete policy loss and entropy loss)
12491249 ppo_discrete_batch = ppo_policy_data (
12501250 output ['logit' ]['action_type' ], data ['logit' ]['action_type' ], data ['action' ]['action_type' ], adv ,
1251- data ['weight' ]
1251+ data ['weight' ], None
12521252 )
12531253 ppo_discrete_loss , ppo_discrete_info = ppo_policy_error (ppo_discrete_batch , self ._clip_ratio )
12541254 # continuous part (continuous policy loss and entropy loss, value loss)
12551255 ppo_continuous_batch = ppo_data (
12561256 output ['logit' ]['action_args' ], data ['logit' ]['action_args' ], data ['action' ]['action_args' ],
1257- output ['value' ], data ['value' ], adv , data ['return' ], data ['weight' ]
1257+ output ['value' ], data ['value' ], adv , data ['return' ], data ['weight' ], None
12581258 )
12591259 ppo_continuous_loss , ppo_continuous_info = ppo_error_continuous (ppo_continuous_batch , self ._clip_ratio )
12601260 # sum discrete and continuous loss
@@ -1279,22 +1279,22 @@ def _forward_learn(self, data: List[Dict[str, Any]]) -> Dict[str, Any]:
12791279
12801280 # Calculate ppo loss
12811281 if self ._action_space == 'continuous' :
1282- ppodata = ppo_policy_data (output ['logit' ], data ['logit' ], data ['action' ], adv , data ['weight' ])
1282+ ppodata = ppo_policy_data (output ['logit' ], data ['logit' ], data ['action' ], adv , data ['weight' ], None )
12831283 ppo_policy_loss , ppo_info = ppo_policy_error_continuous (ppodata , self ._clip_ratio )
12841284 elif self ._action_space == 'discrete' :
1285- ppodata = ppo_policy_data (output ['logit' ], data ['logit' ], data ['action' ], adv , data ['weight' ])
1285+ ppodata = ppo_policy_data (output ['logit' ], data ['logit' ], data ['action' ], adv , data ['weight' ], None )
12861286 ppo_policy_loss , ppo_info = ppo_policy_error (ppodata , self ._clip_ratio )
12871287 elif self ._action_space == 'hybrid' :
12881288 # discrete part (discrete policy loss and entropy loss)
12891289 ppo_discrete_data = ppo_policy_data (
12901290 output ['logit' ]['action_type' ], data ['logit' ]['action_type' ], data ['action' ]['action_type' ], adv ,
1291- data ['weight' ]
1291+ data ['weight' ], None
12921292 )
12931293 ppo_discrete_loss , ppo_discrete_info = ppo_policy_error (ppo_discrete_data , self ._clip_ratio )
12941294 # continuous part (continuous policy loss and entropy loss, value loss)
12951295 ppo_continuous_data = ppo_policy_data (
12961296 output ['logit' ]['action_args' ], data ['logit' ]['action_args' ], data ['action' ]['action_args' ], adv ,
1297- data ['weight' ]
1297+ data ['weight' ], None
12981298 )
12991299 ppo_continuous_loss , ppo_continuous_info = ppo_policy_error_continuous (
13001300 ppo_continuous_data , self ._clip_ratio
@@ -1804,13 +1804,13 @@ def _forward_learn(self, data: Dict[str, Any]) -> Dict[str, Any]:
18041804 if self ._action_space == 'continuous' :
18051805 ppo_batch = ppo_data (
18061806 output ['logit' ], batch ['logit' ], batch ['action' ], output ['value' ], batch ['value' ], adv ,
1807- batch ['return' ], batch ['weight' ]
1807+ batch ['return' ], batch ['weight' ], None
18081808 )
18091809 ppo_loss , ppo_info = ppo_error_continuous (ppo_batch , self ._clip_ratio )
18101810 elif self ._action_space == 'discrete' :
18111811 ppo_batch = ppo_data (
18121812 output ['logit' ], batch ['logit' ], batch ['action' ], output ['value' ], batch ['value' ], adv ,
1813- batch ['return' ], batch ['weight' ]
1813+ batch ['return' ], batch ['weight' ], None
18141814 )
18151815 ppo_loss , ppo_info = ppo_error (ppo_batch , self ._clip_ratio )
18161816
0 commit comments