diff --git a/torchrl/objectives/a2c.py b/torchrl/objectives/a2c.py index 19b64e6111d..96ed0e1ab77 100644 --- a/torchrl/objectives/a2c.py +++ b/torchrl/objectives/a2c.py @@ -558,9 +558,8 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: td_out.set("loss_critic", loss_critic) if value_clip_fraction is not None: td_out.set("value_clip_fraction", value_clip_fraction) - loss_mask = tensordict.get("shifted_valid", default=None) td_out = td_out.named_apply( - lambda name, value: self._reduce_loss(value, mask=loss_mask).squeeze(-1) + lambda name, value: self._reduce_loss(value, tensordict).squeeze(-1) if name.startswith("loss_") else value, ) diff --git a/torchrl/objectives/bc.py b/torchrl/objectives/bc.py index 3d70f37d53d..777b9fb3ea2 100644 --- a/torchrl/objectives/bc.py +++ b/torchrl/objectives/bc.py @@ -17,7 +17,6 @@ from tensordict.utils import NestedKey from torchrl.objectives.common import LossModule -from torchrl.objectives.utils import _reduce class BCLoss(LossModule): @@ -325,10 +324,9 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: "elementwise loss (e.g. loss_function='l1'), not a " "distribution-based (NLL) or cross-entropy loss." ) - while mask.ndim < loss.ndim: - mask = mask.unsqueeze(-1) - mask = mask.expand_as(loss) - loss = _reduce(loss, reduction=self.reduction, mask=mask) + loss = self._reduce_loss(loss, tensordict=tensordict, mask=mask) + if mask is not None and self.reduction == "mean": + loss = loss.masked_fill(~mask.any(), torch.nan) td_out = TensorDict({"loss_bc": loss}) self._clear_weakrefs( diff --git a/torchrl/objectives/common.py b/torchrl/objectives/common.py index 1f32cbced5b..31b56351f20 100644 --- a/torchrl/objectives/common.py +++ b/torchrl/objectives/common.py @@ -279,12 +279,22 @@ def _reduce_loss( ) -> torch.Tensor: if reduction is None: reduction = self.reduction - if mask is None and tensordict is not None: - mask = tensordict.get("shifted_valid", default=None) + preserve_shape = mask is None if mask is not None: mask = self._expand_loss_mask(mask, loss) + if tensordict is not None: + for mask_key in (("collector", "mask"), "shifted_valid"): + tensordict_mask = tensordict.get(mask_key, default=None) + if tensordict_mask is not None: + tensordict_mask = self._expand_loss_mask(tensordict_mask, loss) + mask = tensordict_mask if mask is None else mask & tensordict_mask + if mask is not None: if weights is not None and weights.shape != loss.shape: weights = self._expand_loss_mask(weights, loss) + if reduction == "none" and preserve_shape: + if weights is not None: + loss = loss * weights + return torch.where(mask, loss, torch.zeros_like(loss)) if weights is None and reduction == "mean": return (loss * mask.to(loss.dtype)).sum() / mask.sum().clamp_min(1) if weights is None and reduction == "sum": diff --git a/torchrl/objectives/deprecated.py b/torchrl/objectives/deprecated.py index 0989dfe82e5..f37ee91ea00 100644 --- a/torchrl/objectives/deprecated.py +++ b/torchrl/objectives/deprecated.py @@ -335,6 +335,10 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: raise RuntimeError( f"QVal and actor loss have different shape: {loss_qval.shape} and {loss_actor.shape}" ) + loss_tensordict = tensordict.unsqueeze(0) + loss_actor = self._reduce_loss(loss_actor, loss_tensordict) + loss_qval = self._reduce_loss(loss_qval, loss_tensordict) + loss_alpha = self._reduce_loss(loss_alpha, tensordict) td_out = TensorDict( { "loss_actor": loss_actor, @@ -345,12 +349,6 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: }, [], ) - td_out = td_out.named_apply( - lambda name, value: self._reduce_loss(value, tensordict=tensordict) - if name.startswith("loss_") - else value, - batch_size=[], - ) self._clear_weakrefs( tensordict, td_out, diff --git a/torchrl/objectives/ppo.py b/torchrl/objectives/ppo.py index d3eba0207aa..ee9a8b9f970 100644 --- a/torchrl/objectives/ppo.py +++ b/torchrl/objectives/ppo.py @@ -985,9 +985,8 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: td_out.set("value_clip_fraction", value_clip_fraction) if explained_variance is not None: td_out.set("explained_variance", explained_variance) - loss_mask = tensordict.get("shifted_valid", default=None) td_out = td_out.named_apply( - lambda name, value: self._reduce_loss(value, mask=loss_mask).squeeze(-1) + lambda name, value: self._reduce_loss(value, tensordict).squeeze(-1) if name.startswith("loss_") else value, ) @@ -1435,9 +1434,8 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: ratio = log_weight.exp() td_out.set("max_ratio", ratio.max()) td_out.set("mean_ratio", ratio.mean()) - loss_mask = tensordict.get("shifted_valid", default=None) td_out = td_out.named_apply( - lambda name, value: self._reduce_loss(value, mask=loss_mask).squeeze(-1) + lambda name, value: self._reduce_loss(value, tensordict).squeeze(-1) if name.startswith("loss_") else value, ) @@ -1796,9 +1794,8 @@ def forward(self, tensordict: TensorDictBase) -> TensorDict: td_out.set("value_clip_fraction", value_clip_fraction) if explained_variance is not None: td_out.set("explained_variance", explained_variance) - loss_mask = tensordict_copy.get("shifted_valid", default=None) td_out = td_out.named_apply( - lambda name, value: self._reduce_loss(value, mask=loss_mask).squeeze(-1) + lambda name, value: self._reduce_loss(value, tensordict_copy).squeeze(-1) if name.startswith("loss_") else value, ) diff --git a/torchrl/objectives/redq.py b/torchrl/objectives/redq.py index 78ced67c34d..2ff91819d72 100644 --- a/torchrl/objectives/redq.py +++ b/torchrl/objectives/redq.py @@ -621,7 +621,6 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: "next.state_value": next_state_value.detach(), "target_value": target_value.detach(), } - td_out = TensorDict(out, []) else: out = { "loss_actor": loss_actor, @@ -634,20 +633,17 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: "next.state_value": next_state_value.detach(), "target_value": target_value.detach(), } - td_out = TensorDict(out, []) - if self.reduction != "none": - loss_mask = tensordict.get("shifted_valid", default=None) - if loss_mask is not None: - loss_mask = loss_mask.unsqueeze(0) - td_out = td_out.named_apply( - lambda name, value: self._reduce_loss(value, mask=loss_mask) - if name.startswith("loss_") - else value, - ) - elif self.scalar_output_mode == "non_tensor": - # Move scalars to non-tensor after creation - td_out.set_non_tensor("alpha", td_out.pop("alpha")) - td_out.set_non_tensor("entropy", td_out.pop("entropy")) + td_out = TensorDict(out, []) + loss_tensordict = tensordict.unsqueeze(0) + td_out = td_out.named_apply( + lambda name, value: self._reduce_loss(value, loss_tensordict) + if name.startswith("loss_") + else value, + ) + if self.reduction == "none" and self.scalar_output_mode == "non_tensor": + # Move scalars to non-tensor after creation + td_out.set_non_tensor("alpha", td_out.pop("alpha")) + td_out.set_non_tensor("entropy", td_out.pop("entropy")) self._clear_weakrefs( tensordict, td_out, diff --git a/torchrl/objectives/reinforce.py b/torchrl/objectives/reinforce.py index 7c27d6a7968..e222f5d2081 100644 --- a/torchrl/objectives/reinforce.py +++ b/torchrl/objectives/reinforce.py @@ -393,9 +393,8 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: td_out.set("loss_value", loss_value) if value_clip_fraction is not None: td_out.set("value_clip_fraction", value_clip_fraction) - loss_mask = tensordict.get("shifted_valid", default=None) td_out = td_out.named_apply( - lambda name, value: self._reduce_loss(value, mask=loss_mask).squeeze(-1) + lambda name, value: self._reduce_loss(value, tensordict).squeeze(-1) if name.startswith("loss_") else value, ) diff --git a/torchrl/objectives/rnd.py b/torchrl/objectives/rnd.py index 2acf3086c9c..640206bc940 100644 --- a/torchrl/objectives/rnd.py +++ b/torchrl/objectives/rnd.py @@ -14,7 +14,6 @@ from torchrl.envs.transforms.rnd import RunningMeanStd from torchrl.objectives.common import LossModule -from torchrl.objectives.utils import _reduce class RNDLoss(LossModule): @@ -129,6 +128,7 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: per_sample_loss = (pred_feat - target_feat).pow(2).mean(dim=-1) + mask = None if self.update_fraction < 1.0: mask = ( torch.rand(per_sample_loss.shape, device=per_sample_loss.device) @@ -137,11 +137,8 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: per_sample_loss = torch.where( mask, per_sample_loss, torch.zeros_like(per_sample_loss) ) - if self.reduction == "mean": - loss = per_sample_loss.sum() / mask.sum().clamp_min(1) - else: - loss = _reduce(per_sample_loss, self.reduction) - else: - loss = _reduce(per_sample_loss, self.reduction) + if self.reduction == "none": + mask = None + loss = self._reduce_loss(per_sample_loss, tensordict=tensordict, mask=mask) return TensorDict({"loss_predictor": loss}, batch_size=[])