@@ -751,23 +751,23 @@ def get_actions(
751751 dtype = torch .uint8 if (self .discrete or self .multibinary ) else torch .float32
752752 return actions .to (dtype = dtype ), hidden_state
753753
754- def _critic_ensemble_to_td_target (self , ensemble_td_target : torch .Tensor ):
755- B , L , C , G , _ = ensemble_td_target .shape
754+ def _reduce_critic_ensemble (self , q_ensemble : torch .Tensor ) -> torch . Tensor :
755+ B , L , C , G , _ = q_ensemble .shape
756756 # random subset of critic ensemble
757757 random_subset = torch .randint (
758758 low = 0 ,
759759 high = C ,
760760 size = (B , L , self .num_critics_td , G , 1 ),
761- device = ensemble_td_target .device ,
761+ device = q_ensemble .device ,
762762 )
763- td_target_rand = torch .take_along_dim (ensemble_td_target , random_subset , dim = 2 )
763+ q_subset = torch .take_along_dim (q_ensemble , random_subset , dim = 2 )
764764 if self .online_coeff > 0 :
765765 # clipped double q
766- td_target = td_target_rand .min (2 , keepdims = True ).values
766+ q_reduced = q_subset .min (2 , keepdims = True ).values
767767 else :
768768 # without DPG updates the usual min creates strong underestimation. take mean instead
769- td_target = td_target_rand .mean (2 , keepdims = True )
770- return td_target
769+ q_reduced = q_subset .mean (2 , keepdims = True )
770+ return q_reduced
771771
772772 def _compute_loss (
773773 self ,
@@ -897,7 +897,7 @@ def forward(self, batch: Batch, log_step: bool) -> torch.Tensor:
897897 # Q_target(s', a')
898898 q_targ_sp_ap_gp = self .popart (self .target_critics (* sp_ap_gp ).mean (0 ), normalized = False )
899899 assert q_targ_sp_ap_gp .shape == (B , L - 1 , C , G , 1 )
900- q_reduced = self ._critic_ensemble_to_td_target (q_targ_sp_ap_gp )
900+ q_reduced = self ._reduce_critic_ensemble (q_targ_sp_ap_gp )
901901 assert q_reduced .shape == (B , L - 1 , 1 , G , 1 )
902902 nstep_mask = state_mask .float ().unsqueeze (- 1 ).unsqueeze (- 1 ) # (B, L-1, 1, 1, 1)
903903 td_target = self ._nstep_fn (r , d , q_reduced , gamma , mask = nstep_mask )
@@ -1258,7 +1258,7 @@ def forward(self, batch: Batch, log_step: bool):
12581258 assert q_targ_sp_ap_gp .probs .shape == (K_c , B , L - 1 , C , G , Bins )
12591259 q_targ_sp_ap_gp = self .target_critics .bin_dist_to_raw_vals (q_targ_sp_ap_gp ).mean (0 )
12601260 assert q_targ_sp_ap_gp .shape == (B , L - 1 , C , G , 1 )
1261- q_reduced = self ._critic_ensemble_to_td_target (q_targ_sp_ap_gp )
1261+ q_reduced = self ._reduce_critic_ensemble (q_targ_sp_ap_gp )
12621262 assert q_reduced .shape == (B , L - 1 , 1 , G , 1 )
12631263 nstep_mask = state_mask .float ().unsqueeze (- 1 ).unsqueeze (- 1 ) # (B, L-1, 1, 1, 1)
12641264 td_target = self ._nstep_fn (r , d , q_reduced , gamma , mask = nstep_mask )
0 commit comments