From 4c9a62ba2e5253a8c048412e395e6a880fe13d0d Mon Sep 17 00:00:00 2001 From: gtnv <297143481+gtnv@users.noreply.github.com> Date: Sat, 1 Aug 2026 14:10:44 -1000 Subject: [PATCH 1/5] fix/loss-mask PR #3888 routed objectives through _reduce_loss, but the helper still ignored ("collector", "mask") because the corresponding change in #3850 never merged. This fix excludes padded timesteps from loss reduction so they do not affect the loss or gradients. --- torchrl/objectives/a2c.py | 3 +-- torchrl/objectives/bc.py | 8 +++---- torchrl/objectives/common.py | 26 ++++++++++++++++++-- torchrl/objectives/decision_transformer.py | 24 ++++++++++--------- torchrl/objectives/deprecated.py | 10 ++++---- torchrl/objectives/ppo.py | 9 +++---- torchrl/objectives/redq.py | 28 ++++++++++------------ torchrl/objectives/reinforce.py | 3 +-- torchrl/objectives/rnd.py | 11 ++++----- 9 files changed, 66 insertions(+), 56 deletions(-) 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..b5ec9cb479d 100644 --- a/torchrl/objectives/common.py +++ b/torchrl/objectives/common.py @@ -264,6 +264,8 @@ def _set_deprecated_ctor_keys(self, **kwargs) -> None: @staticmethod def _expand_loss_mask(mask: torch.Tensor, loss: torch.Tensor) -> torch.Tensor: + while mask.ndim > loss.ndim and mask.shape[-1] == 1: + mask = mask.squeeze(-1) if mask.ndim < loss.ndim: mask = mask.reshape(mask.shape + (1,) * (loss.ndim - mask.ndim)) return mask.expand_as(loss) @@ -276,15 +278,35 @@ def _reduce_loss( mask: torch.Tensor | None = None, reduction: str | None = None, weights: torch.Tensor | None = None, + preserve_shape: bool | None = None, ) -> 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) + infer_preserve_shape = preserve_shape is None + if infer_preserve_shape: + 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: + if mask_key == "shifted_valid" and infer_preserve_shape: + preserve_shape = False + 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 not None and reduction == "mean": + masked_weights = weights * mask.to(weights.dtype) + denominator = masked_weights.sum() + denominator = denominator.masked_fill(denominator == 0, 1) + return (loss * masked_weights).sum() / denominator 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/decision_transformer.py b/torchrl/objectives/decision_transformer.py index e74e9c9f50b..5f579da9cc8 100644 --- a/torchrl/objectives/decision_transformer.py +++ b/torchrl/objectives/decision_transformer.py @@ -255,24 +255,26 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: "loss_alpha": loss_alpha, } - # Handle batch_size and scalar values (alpha, entropy) based on reduction mode + # Handle batch_size and scalar values based on reduction mode if self.reduction == "none": batch_size = tensordict.batch_size td_out = TensorDict(out, batch_size=batch_size) - if self.scalar_output_mode == "non_tensor": - td_out.set_non_tensor("alpha", self.alpha.detach()) - td_out.set_non_tensor("entropy", entropy.detach().mean()) else: out["entropy"] = entropy.detach().mean() out["alpha"] = self.alpha.detach() td_out = TensorDict(out, []) - td_out = td_out.named_apply( - lambda name, value: self._reduce_loss( - value, tensordict=tensordict - ).squeeze(-1) - if name.startswith("loss_") - else value, - ) + td_out = td_out.named_apply( + lambda name, value: self._reduce_loss( + value, + tensordict, + preserve_shape=self.reduction == "none", + ).squeeze(-1) + if name.startswith("loss_") + else value, + ) + if self.reduction == "none" and self.scalar_output_mode == "non_tensor": + td_out.set_non_tensor("alpha", self.alpha.detach()) + td_out.set_non_tensor("entropy", entropy.detach().mean()) self._clear_weakrefs( tensordict, td_out, 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..c407dd15015 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,19 @@ 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, tensordict=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=[]) From c80ee03279298ac810f08fa979b90f3451c02980 Mon Sep 17 00:00:00 2001 From: gtnv <297143481+gtnv@users.noreply.github.com> Date: Sat, 1 Aug 2026 18:08:20 -1000 Subject: [PATCH 2/5] fix squeeze syntax & REDQloss func --- torchrl/objectives/decision_transformer.py | 2 +- torchrl/objectives/redq.py | 17 +++++++++-------- 2 files changed, 10 insertions(+), 9 deletions(-) diff --git a/torchrl/objectives/decision_transformer.py b/torchrl/objectives/decision_transformer.py index 5f579da9cc8..48d2d587cfa 100644 --- a/torchrl/objectives/decision_transformer.py +++ b/torchrl/objectives/decision_transformer.py @@ -268,7 +268,7 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: value, tensordict, preserve_shape=self.reduction == "none", - ).squeeze(-1) + ) if name.startswith("loss_") else value, ) diff --git a/torchrl/objectives/redq.py b/torchrl/objectives/redq.py index c407dd15015..950675064f0 100644 --- a/torchrl/objectives/redq.py +++ b/torchrl/objectives/redq.py @@ -634,14 +634,15 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: "target_value": target_value.detach(), } td_out = TensorDict(out, []) - loss_tensordict = tensordict.unsqueeze(0) - td_out = td_out.named_apply( - lambda name, value: ( - self._reduce_loss(value, tensordict=loss_tensordict) - if name.startswith("loss_") - else value - ), - ) + if self.reduction != "none": + loss_tensordict = tensordict.unsqueeze(0) + td_out = td_out.named_apply( + lambda name, value: ( + self._reduce_loss(value, tensordict=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")) From 38adc30ffb115ad84a14978cc4b2c2f3f6ebc0b9 Mon Sep 17 00:00:00 2001 From: gtnv <297143481+gtnv@users.noreply.github.com> Date: Sat, 1 Aug 2026 19:37:52 -1000 Subject: [PATCH 3/5] keep current weight behavior same --- torchrl/objectives/common.py | 5 ----- 1 file changed, 5 deletions(-) diff --git a/torchrl/objectives/common.py b/torchrl/objectives/common.py index b5ec9cb479d..c90db3dbb94 100644 --- a/torchrl/objectives/common.py +++ b/torchrl/objectives/common.py @@ -302,11 +302,6 @@ def _reduce_loss( if weights is not None: loss = loss * weights return torch.where(mask, loss, torch.zeros_like(loss)) - if weights is not None and reduction == "mean": - masked_weights = weights * mask.to(weights.dtype) - denominator = masked_weights.sum() - denominator = denominator.masked_fill(denominator == 0, 1) - return (loss * masked_weights).sum() / denominator 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": From f5993b28c55f38f538a2fa51fd9ef5cbf8ed7e42 Mon Sep 17 00:00:00 2001 From: gtnv <297143481+gtnv@users.noreply.github.com> Date: Sat, 1 Aug 2026 20:56:46 -1000 Subject: [PATCH 4/5] update collector masks --- torchrl/objectives/common.py | 7 ++----- torchrl/objectives/decision_transformer.py | 22 ++++++++++------------ torchrl/objectives/redq.py | 8 +++----- 3 files changed, 15 insertions(+), 22 deletions(-) diff --git a/torchrl/objectives/common.py b/torchrl/objectives/common.py index c90db3dbb94..bc69605204f 100644 --- a/torchrl/objectives/common.py +++ b/torchrl/objectives/common.py @@ -278,20 +278,17 @@ def _reduce_loss( mask: torch.Tensor | None = None, reduction: str | None = None, weights: torch.Tensor | None = None, - preserve_shape: bool | None = None, ) -> torch.Tensor: if reduction is None: reduction = self.reduction - infer_preserve_shape = preserve_shape is None - if infer_preserve_shape: - preserve_shape = mask is 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: - if mask_key == "shifted_valid" and infer_preserve_shape: + if mask_key == "shifted_valid": preserve_shape = False tensordict_mask = self._expand_loss_mask(tensordict_mask, loss) mask = tensordict_mask if mask is None else mask & tensordict_mask diff --git a/torchrl/objectives/decision_transformer.py b/torchrl/objectives/decision_transformer.py index 48d2d587cfa..e74e9c9f50b 100644 --- a/torchrl/objectives/decision_transformer.py +++ b/torchrl/objectives/decision_transformer.py @@ -255,26 +255,24 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: "loss_alpha": loss_alpha, } - # Handle batch_size and scalar values based on reduction mode + # Handle batch_size and scalar values (alpha, entropy) based on reduction mode if self.reduction == "none": batch_size = tensordict.batch_size td_out = TensorDict(out, batch_size=batch_size) + if self.scalar_output_mode == "non_tensor": + td_out.set_non_tensor("alpha", self.alpha.detach()) + td_out.set_non_tensor("entropy", entropy.detach().mean()) else: out["entropy"] = entropy.detach().mean() out["alpha"] = self.alpha.detach() td_out = TensorDict(out, []) - td_out = td_out.named_apply( - lambda name, value: self._reduce_loss( - value, - tensordict, - preserve_shape=self.reduction == "none", + td_out = td_out.named_apply( + lambda name, value: self._reduce_loss( + value, tensordict=tensordict + ).squeeze(-1) + if name.startswith("loss_") + else value, ) - if name.startswith("loss_") - else value, - ) - if self.reduction == "none" and self.scalar_output_mode == "non_tensor": - td_out.set_non_tensor("alpha", self.alpha.detach()) - td_out.set_non_tensor("entropy", entropy.detach().mean()) self._clear_weakrefs( tensordict, td_out, diff --git a/torchrl/objectives/redq.py b/torchrl/objectives/redq.py index 950675064f0..8b850960232 100644 --- a/torchrl/objectives/redq.py +++ b/torchrl/objectives/redq.py @@ -637,11 +637,9 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: if self.reduction != "none": loss_tensordict = tensordict.unsqueeze(0) td_out = td_out.named_apply( - lambda name, value: ( - self._reduce_loss(value, tensordict=loss_tensordict) - if name.startswith("loss_") - else value - ), + 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 From abb96df8ebfb6f814ce4d609362522a1a4c762d9 Mon Sep 17 00:00:00 2001 From: gtnv <297143481+gtnv@users.noreply.github.com> Date: Sat, 1 Aug 2026 23:21:32 -1000 Subject: [PATCH 5/5] feedback/reduction value consistency Unify unreduced mask behavior. Apply collector masks to REDQ under reduction="none". Remove unreachable mask-shape handling. --- torchrl/objectives/common.py | 4 ---- torchrl/objectives/redq.py | 13 ++++++------- 2 files changed, 6 insertions(+), 11 deletions(-) diff --git a/torchrl/objectives/common.py b/torchrl/objectives/common.py index bc69605204f..31b56351f20 100644 --- a/torchrl/objectives/common.py +++ b/torchrl/objectives/common.py @@ -264,8 +264,6 @@ def _set_deprecated_ctor_keys(self, **kwargs) -> None: @staticmethod def _expand_loss_mask(mask: torch.Tensor, loss: torch.Tensor) -> torch.Tensor: - while mask.ndim > loss.ndim and mask.shape[-1] == 1: - mask = mask.squeeze(-1) if mask.ndim < loss.ndim: mask = mask.reshape(mask.shape + (1,) * (loss.ndim - mask.ndim)) return mask.expand_as(loss) @@ -288,8 +286,6 @@ def _reduce_loss( for mask_key in (("collector", "mask"), "shifted_valid"): tensordict_mask = tensordict.get(mask_key, default=None) if tensordict_mask is not None: - if mask_key == "shifted_valid": - preserve_shape = False 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: diff --git a/torchrl/objectives/redq.py b/torchrl/objectives/redq.py index 8b850960232..2ff91819d72 100644 --- a/torchrl/objectives/redq.py +++ b/torchrl/objectives/redq.py @@ -634,13 +634,12 @@ def forward(self, tensordict: TensorDictBase) -> TensorDictBase: "target_value": target_value.detach(), } td_out = TensorDict(out, []) - if self.reduction != "none": - 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, - ) + 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"))