From 911019d2a84a04579cd5f2cbe9d30e87da0c3f30 Mon Sep 17 00:00:00 2001 From: HuiyingLi Date: Sun, 30 Aug 2026 11:48:26 -0700 Subject: [PATCH] fix(loss): keep autograd and honor reduction in NemotronParseLoss An all-ignored batch returned a graph-free torch.tensor(0.0), breaking backward; reduction='sum' silently divided by valid_tokens anyway. Return loss_sum * 0 for the empty case and dispatch reduction explicitly. Co-Authored-By: Claude Fable 5 Signed-off-by: HuiyingLi --- .../nemotron_parse/nemotron_parse_loss.py | 11 ++++-- .../loss/test_nemotron_parse_loss.py | 35 +++++++++++++++---- 2 files changed, 37 insertions(+), 9 deletions(-) diff --git a/nemo_automodel/components/models/nemotron_parse/nemotron_parse_loss.py b/nemo_automodel/components/models/nemotron_parse/nemotron_parse_loss.py index e36da4bc1c..261dac8b83 100644 --- a/nemo_automodel/components/models/nemotron_parse/nemotron_parse_loss.py +++ b/nemo_automodel/components/models/nemotron_parse/nemotron_parse_loss.py @@ -105,13 +105,18 @@ def forward( loss_full[coordinate_mask] *= self.coordinate_weight valid_tokens = (labels != self.ignore_index).sum() + loss_sum = loss_full.sum() if valid_tokens == 0: - return torch.tensor(0.0, device=logits.device, dtype=logits.dtype) + return loss_sum * 0 if num_label_tokens is not None: assert self.reduction == "sum", ( f"num_label_tokens is only supported when reduction='sum', got reduction='{self.reduction}'" ) - return loss_full.sum() / (num_label_tokens + 1e-6) + return loss_sum / (num_label_tokens + 1e-6) - return loss_full.sum() / (valid_tokens + 1e-6) + if self.reduction == "sum": + return loss_sum + if self.reduction == "mean": + return loss_sum / (valid_tokens + 1e-6) + raise ValueError(f"Unsupported reduction={self.reduction!r}; expected 'sum' or 'mean'") diff --git a/tests/unit_tests/loss/test_nemotron_parse_loss.py b/tests/unit_tests/loss/test_nemotron_parse_loss.py index 1b89569ce9..00253b4476 100644 --- a/tests/unit_tests/loss/test_nemotron_parse_loss.py +++ b/tests/unit_tests/loss/test_nemotron_parse_loss.py @@ -76,7 +76,7 @@ def test_backward_compatibility(): logits = torch.randn(2, 10, 100) labels = torch.randint(0, 100, (2, 10)) - loss_fn = NemotronParseLoss(coordinate_weight=10.0, class_token_start_idx=50000) + loss_fn = NemotronParseLoss(coordinate_weight=10.0, class_token_start_idx=50000, reduction="mean") loss_new = loss_fn(logits=logits, labels=labels) loss_ref = _compute_reference_loss(logits, labels) @@ -126,7 +126,7 @@ def test_fp32_upcast(): assert torch.isfinite(loss_fp32) assert torch.isfinite(loss_bf16) - assert torch.allclose(loss_fp32, loss_bf16, rtol=1e-2) + assert torch.allclose(loss_fp32, loss_bf16.float(), rtol=1e-2) def test_invalid_logits_shape(): @@ -265,10 +265,12 @@ def test_mixed_coordinate_and_regular_tokens(): """Test loss computation with mixed coordinate and regular tokens.""" torch.manual_seed(42) logits = torch.randn(2, 10, 50150) # Increased vocab size to accommodate labels - labels = torch.tensor([ - [10, 20, 50001, 50002, 30, 40, 50003, 50, 60, 70], # Mixed - [50100, 50101, 50102, 100, 200, 300, 50103, 400, 500, 600] # Mixed - ]) + labels = torch.tensor( + [ + [10, 20, 50001, 50002, 30, 40, 50003, 50, 60, 70], # Mixed + [50100, 50101, 50102, 100, 200, 300, 50103, 400, 500, 600], # Mixed + ] + ) loss_fn = NemotronParseLoss(coordinate_weight=10.0, class_token_start_idx=50000) loss = loss_fn(logits=logits, labels=labels) @@ -346,3 +348,24 @@ def test_reduction_parameter_stored(): assert loss_fn_sum.reduction == "sum" assert loss_fn_mean.reduction == "mean" + + +def test_sum_and_mean_reductions_have_distinct_normalization(): + logits = torch.randn(2, 5, 100) + labels = torch.randint(0, 100, (2, 5)) + labels[0, :2] = -100 + loss_sum = NemotronParseLoss(reduction="sum")(logits=logits, labels=labels) + loss_mean = NemotronParseLoss(reduction="mean")(logits=logits, labels=labels) + + torch.testing.assert_close(loss_sum, loss_mean * (labels != -100).sum()) + + +def test_all_ignored_sum_keeps_a_zero_autograd_path(): + logits = torch.randn(2, 5, 100, requires_grad=True) + labels = torch.full((2, 5), -100) + + loss = NemotronParseLoss(reduction="sum")(logits=logits, labels=labels) + loss.backward() + + assert logits.grad is not None + assert torch.count_nonzero(logits.grad) == 0