Skip to content

Commit 4be39df

Browse files
committed
rename td target ensemble method
1 parent d0d2ad5 commit 4be39df

1 file changed

Lines changed: 9 additions & 9 deletions

File tree

amago/agent.py

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)