diff --git a/tzrec/loss/listwise_rank_loss.py b/tzrec/loss/listwise_rank_loss.py index c2c67e59..002c4b66 100644 --- a/tzrec/loss/listwise_rank_loss.py +++ b/tzrec/loss/listwise_rank_loss.py @@ -12,7 +12,6 @@ """Per-request list-wise InfoNCE over a jagged candidate list (paper 2.5).""" import math -from typing import Optional import torch from torch import nn @@ -88,7 +87,6 @@ def forward( logits: torch.Tensor, labels: torch.Tensor, lengths: torch.Tensor, - loss_weight: Optional[torch.Tensor] = None, ) -> torch.Tensor: """Compute the list-wise InfoNCE term. @@ -102,8 +100,6 @@ def forward( non-zero value counts as a positive. lengths (torch.Tensor): ``(B,)`` candidates per request, summing to ``total``. - loss_weight (torch.Tensor, optional): scalar multiplier, used - to turn the local-batch mean into a global-batch mean. Returns: torch.Tensor: scalar loss -- the mean of per-request losses over @@ -138,14 +134,11 @@ def forward( per_request = torch.nan_to_num(per_request, nan=0.0) # Divide by the *total* request count, not the valid count: the - # `enable_global_average_loss` rescale in ``RankModel._loss_impl`` + # `enable_global_average_loss` rescale in ``DlrmHSTU.loss`` # divides by the total count too, so the two denominators agree # and DDP's cross-rank gradient average stays an unbiased global # mean even when the masked-out fraction differs across ranks. # The max() keeps an empty (zero-request) batch at 0 instead of # the NaN of 0/0, which would poison every parameter through the # all-reduced gradient. - loss = (per_request * valid).sum() / fx_size0_max1(lengths) - if loss_weight is not None: - loss = loss * loss_weight - return loss + return (per_request * valid).sum() / fx_size0_max1(lengths) diff --git a/tzrec/loss/listwise_rank_loss_test.py b/tzrec/loss/listwise_rank_loss_test.py index 5bef46b1..c8263fd7 100644 --- a/tzrec/loss/listwise_rank_loss_test.py +++ b/tzrec/loss/listwise_rank_loss_test.py @@ -236,18 +236,6 @@ def test_non_positive_temperature_raises(self) -> None: with self.assertRaisesRegex(ValueError, "temperature_init"): ListwiseRankLoss(temperature_init=bad) - def test_loss_weight_is_a_scalar_multiplier(self) -> None: - """``loss_weight`` carries the global-average-loss rescaling.""" - module = ListwiseRankLoss(temperature_init=1.0, learnable_temperature=False) - lengths = torch.tensor([3, 4], dtype=torch.int64) - labels = torch.tensor([1.0, 0, 0, 0, 1, 0, 0]) - logits = torch.randn(7) - - torch.testing.assert_close( - module(logits, labels, lengths, torch.tensor(2.5)), - module(logits, labels, lengths) * 2.5, - ) - def test_integer_labels_are_accepted(self) -> None: """Labels arrive as ints from the bitmask decode in some configs.""" module = ListwiseRankLoss(temperature_init=1.0, learnable_temperature=False) diff --git a/tzrec/models/dlrm_hstu.py b/tzrec/models/dlrm_hstu.py index 32ad7914..14a4b343 100644 --- a/tzrec/models/dlrm_hstu.py +++ b/tzrec/models/dlrm_hstu.py @@ -36,13 +36,13 @@ from tzrec.protos.tower_pb2 import FusionSubTaskConfig from tzrec.utils.config_util import config_to_kwargs from tzrec.utils.fx_util import ( - fx_avg_batch_size, + fx_avg_counts, fx_flip_tensor_dict, fx_int_item, fx_numel, ) -torch.fx.wrap(fx_avg_batch_size) +torch.fx.wrap(fx_avg_counts) torch.fx.wrap(fx_flip_tensor_dict) torch.fx.wrap(fx_int_item) torch.fx.wrap(fx_numel) @@ -283,15 +283,28 @@ def loss( ) -> Dict[str, torch.Tensor]: """Compute loss of the model.""" losses = {} + request_avg_weight = None + sample_avg_weight = None + if self._model_config.enable_global_average_loss: + # Cost-based batching makes both counts ragged across ranks, and + # every task's label spans the same candidates, so one reduction + # of the two serves every loss of every task. + lengths = predictions[TARGET_REPEAT_INTERLEAVE_KEY] + avg_counts = fx_avg_counts(lengths) + request_avg_weight = lengths.size(0) / avg_counts[0] + sample_avg_weight = lengths.sum() / avg_counts[1] + for task_cfg in self._task_configs: task_name = task_cfg.task_name label = self._get_label(batch, task_cfg) - loss_weight = None - if self._model_config.enable_global_average_loss: - avg_batch_size = fx_avg_batch_size(label) - loss_weight = label.size(0) / avg_batch_size - for loss_cfg in task_cfg.losses: + # The list-wise term reduces over requests, the rest over + # candidates, so they take different rescaling factors. + loss_weight = ( + request_avg_weight + if loss_cfg.WhichOneof("loss") == "listwise_rank_loss" + else sample_avg_weight + ) task_losses = self._loss_impl( predictions, batch, diff --git a/tzrec/models/dlrm_hstu_onerank_test.py b/tzrec/models/dlrm_hstu_onerank_test.py index 3410300b..dc356d00 100644 --- a/tzrec/models/dlrm_hstu_onerank_test.py +++ b/tzrec/models/dlrm_hstu_onerank_test.py @@ -20,8 +20,10 @@ import unittest from typing import List, Optional +from unittest import mock import torch +import torch.distributed as dist from hypothesis import Verbosity, given from hypothesis import strategies as st from torchrec import JaggedTensor, KeyedJaggedTensor @@ -42,6 +44,7 @@ tower_pb2, ) from tzrec.protos.models import multi_task_rank_pb2 +from tzrec.utils import fx_util from tzrec.utils.state_dict_util import init_parameters from tzrec.utils.test_util import ( TestGraphType, @@ -648,6 +651,70 @@ def test_listwise_loss_is_scaled_by_loss_weight(self) -> None: base_losses[f"binary_cross_entropy_{task_name}"], ) + @unittest.skipIf(*gpu_unavailable) + def test_global_average_loss_rescales_each_term_by_its_own_denominator( + self, + ) -> None: + """Each loss family takes its own global-average-loss factor. + + The list-wise term is a mean over requests and the point-wise terms + are means over candidates, so each must be rescaled by the ratio of + the count it divided by -- mixing them up leaves the loss finite and + the sign correct, so only a numeric check catches it. Both counts + are ragged across ranks under cost-based batching. + """ + device = torch.device("cuda") + base = _build_model(device=device, listwise_loss=_listwise_loss_cfg()) + scaled = _build_model( + device=device, + listwise_loss=_listwise_loss_cfg(), + enable_global_average_loss=True, + ) + scaled.load_state_dict(base.state_dict()) + for model in (base, scaled): + model.set_kernel(Kernel.PYTORCH) + model.init_loss() + model.eval() + + batch = _build_batch(device=device) + with torch.no_grad(): + predictions = base.predict(batch) + base_losses = base.loss(predictions, batch) + # Emulate one peer rank holding 6 requests / 2 candidates against + # this rank's 2 / 6, so the two ratios differ and neither is 1.0. + peer = torch.tensor([6.0, 2.0], device=device) + with mock.patch.object(fx_util, "dist") as dist_mock: + dist_mock.is_initialized.return_value = True + dist_mock.ReduceOp.AVG = dist.ReduceOp.AVG + dist_mock.all_reduce.side_effect = lambda outcome, op: outcome.copy_( + (outcome + peer) / 2 + ) + scaled_losses = scaled.loss(predictions, batch) + # The flag-off model must take no factor even with a live + # process group, which pins the `global_average` gate: every + # single-process factor is 1.0, so nothing else would catch + # the gate being dropped. + gated_losses = base.loss(predictions, batch) + + # Both axes travel in one reduction, whatever the task or loss count. + self.assertEqual(dist_mock.all_reduce.call_count, 1) + + request_ratio = len(_NUM_TARGETS) / ((len(_NUM_TARGETS) + 6.0) / 2) + candidate_ratio = _TOTAL_TARGETS / ((_TOTAL_TARGETS + 2.0) / 2) + self.assertNotAlmostEqual(request_ratio, candidate_ratio) + + for name, value in base_losses.items(): + torch.testing.assert_close(gated_losses[name], value, msg=name) + + key = "listwise_rank_loss_is_click" + self.assertGreater(base_losses[key].item(), 0.0) + torch.testing.assert_close(scaled_losses[key], base_losses[key] * request_ratio) + for task_name in _TASK_NAMES: + pointwise = f"binary_cross_entropy_{task_name}" + torch.testing.assert_close( + scaled_losses[pointwise], base_losses[pointwise] * candidate_ratio + ) + @unittest.skipIf(*gpu_unavailable) def test_listwise_loss_matches_hand_decoded_reference(self) -> None: """The list-wise term must see decoded labels and its own logits. diff --git a/tzrec/models/rank_model.py b/tzrec/models/rank_model.py index 1dede50d..5280076a 100644 --- a/tzrec/models/rank_model.py +++ b/tzrec/models/rank_model.py @@ -36,9 +36,6 @@ from tzrec.protos.loss_pb2 import LossConfig from tzrec.protos.metric_pb2 import MetricConfig, TrainMetricConfig from tzrec.utils.config_util import config_to_kwargs -from tzrec.utils.fx_util import fx_avg_batch_size - -torch.fx.wrap(fx_avg_batch_size) @torch.fx.wrap @@ -283,21 +280,10 @@ def _loss_impl( "(predictions[TARGET_REPEAT_INTERLEAVE_KEY]), which " "this model does not publish." ) - # The module averages over this rank's request count; rescale - # by the local/global request-count ratio so that DDP's - # cross-rank gradient average comes out as a global mean on a - # ragged batch. Both denominators are total counts, so the - # average stays unbiased even when the masked-out fraction - # differs across ranks. - global_avg_weight = None - if getattr(self._base_model_config, "enable_global_average_loss", False): - global_avg_weight = lengths.size(0) / fx_avg_batch_size(lengths) - losses[loss_name] = self._loss_modules[loss_name]( - pred, label, lengths, global_avg_weight - ) - # The caller's per-candidate loss_weight does not apply here: it - # is sized off the candidate count, not the request count. - loss_weight = None + # NOTE: this loss is a mean over requests, so a loss_weight + # reaching the tail below must be request-level too, not the + # per-candidate weight the sibling losses take. + losses[loss_name] = self._loss_modules[loss_name](pred, label, lengths) else: raise ValueError(f"loss[{loss_type}] is not supported yet.") if loss_weight is not None: diff --git a/tzrec/utils/fx_util.py b/tzrec/utils/fx_util.py index 7b6fdd63..e0f66a3c 100644 --- a/tzrec/utils/fx_util.py +++ b/tzrec/utils/fx_util.py @@ -152,16 +152,31 @@ def fx_numel(x: torch.Tensor) -> int: @torch.fx.wrap -def fx_avg_batch_size(x: torch.Tensor) -> torch.Tensor: - """Fx trace wrapper for the DDP-averaged first dimension of ``x``. +def fx_avg_counts(lengths: torch.Tensor) -> torch.Tensor: + """Fx trace wrapper for the DDP-averaged ``(requests, candidates)`` counts. Used to rescale a local-batch mean loss into a global-batch mean so - DDP's cross-rank gradient average stays unbiased on ragged batches. + DDP's cross-rank gradient average stays unbiased on ragged batches. A + loss that reduces over requests takes the first element and one that + reduces over candidates the second; both travel in one all-reduce + because a scalar collective costs a rank synchronization, not + bandwidth. + + Args: + lengths (torch.Tensor): ``(B,)`` candidates per request. + + Returns: + torch.Tensor: ``(2,)`` cross-rank mean of ``B`` and ``sum(lengths)``. """ - batch_size = torch.tensor(x.size(0), dtype=torch.float, device=x.device) + counts = torch.stack( + [ + torch.tensor(lengths.size(0), dtype=torch.float, device=lengths.device), + lengths.sum().to(torch.float), + ] + ) if dist.is_initialized(): - dist.all_reduce(batch_size, op=dist.ReduceOp.AVG) - return batch_size + dist.all_reduce(counts, op=dist.ReduceOp.AVG) + return counts @torch.fx.wrap diff --git a/tzrec/utils/fx_util_test.py b/tzrec/utils/fx_util_test.py index ee462cc2..9929a763 100644 --- a/tzrec/utils/fx_util_test.py +++ b/tzrec/utils/fx_util_test.py @@ -11,10 +11,10 @@ """Unit tests for ``tzrec.utils.fx_util``. -``fx_avg_batch_size`` was promoted here from a ``DlrmHSTU`` private helper -when ``listwise_rank_loss`` was generalized to ``RankModel``, so every rank -model now consumes it. Its distributed branch (``dist.all_reduce``) is -unreachable from single-process unittests, so it is pinned through the same +``fx_avg_counts`` carries the local/global rescaling factors behind +``enable_global_average_loss``, and ``DlrmHSTU.loss`` is its only caller. +Its distributed branch (``dist.all_reduce``) is unreachable from +single-process unittests, so it is pinned through the same ``mock.patch.object`` pattern as ``tzrec.utils.predict_util_test``. """ @@ -34,7 +34,7 @@ from torchrec.quant.embedding_modules import _permute_kjt from tzrec.utils import fx_util -from tzrec.utils.fx_util import fx_avg_batch_size +from tzrec.utils.fx_util import fx_avg_counts from tzrec.utils.test_util import ( gpu_unavailable, make_test_dir, @@ -214,44 +214,45 @@ def test_torch_export_preserves_optional_weights(self, weighted) -> None: torch.testing.assert_close(result, reference) -class FxAvgBatchSizeTest(unittest.TestCase): - """The local/global rescaling factor of ``enable_global_average_loss``.""" +class FxAvgCountsTest(unittest.TestCase): + """The local/global rescaling factors of ``enable_global_average_loss``.""" - def test_single_process_returns_local_size(self) -> None: - """Without an initialized process group the local size is the answer.""" + def test_single_process_returns_local_counts(self) -> None: + """Without an initialized process group the local counts are the answer.""" if dist.is_initialized(): self.skipTest("dist already initialized in this process") - x = torch.zeros(7) - out = fx_avg_batch_size(x) - self.assertEqual(out.item(), 7.0) + lengths = torch.tensor([3, 0, 4], dtype=torch.int64) + out = fx_avg_counts(lengths) + self.assertEqual(out.tolist(), [3.0, 7.0]) self.assertEqual(out.dtype, torch.float32) - self.assertEqual(out.device, x.device) + self.assertEqual(out.device, lengths.device) def test_empty_shard_is_reported_verbatim(self) -> None: - """A rank with an empty shard contributes 0 to the average.""" - out = fx_avg_batch_size(torch.zeros(0)) - self.assertEqual(out.item(), 0.0) + """A rank with an empty shard contributes 0 to both averages.""" + out = fx_avg_counts(torch.zeros(0, dtype=torch.int64)) + self.assertEqual(out.tolist(), [0.0, 0.0]) - def test_dist_branch_averages_across_ranks(self) -> None: - """The distributed branch must reduce with AVG semantics. + def test_dist_branch_averages_both_axes_in_one_reduction(self) -> None: + """The distributed branch must reduce both counts with AVG semantics. - Two ranks with ragged shard sizes 3 and 5 must both read 4.0 as - the global average; a SUM reduction would read 8.0 and double the - ``local / global`` loss rescaling factor -- silently, because the - loss stays finite and the sign stays correct. + Two ranks holding (3 requests, 10 candidates) and (5, 30) must read + (4.0, 20.0). A SUM reduction would read double and halve the + ``local / global`` loss rescaling factors -- silently, because the + loss stays finite and the sign stays correct. The two axes share one + collective, so a per-axis reduction would show up as a second call. """ - x = torch.zeros(3) + lengths = torch.tensor([4, 6], dtype=torch.int64) + peer = torch.tensor([5.0, 30.0]) with mock.patch.object(fx_util, "dist") as dist_mock: dist_mock.is_initialized.return_value = True # Hand the mock the real enum so the recorded call is assertable. dist_mock.ReduceOp.AVG = dist.ReduceOp.AVG - # Emulate ReduceOp.AVG against a peer that owns 5 rows. - dist_mock.all_reduce.side_effect = lambda outcome, op: outcome.fill_( - (outcome.item() + 5) / 2 + dist_mock.all_reduce.side_effect = lambda outcome, op: outcome.copy_( + (outcome + peer) / 2 ) - out = fx_avg_batch_size(x) + out = fx_avg_counts(lengths) - self.assertEqual(out.item(), 4.0) + self.assertEqual(out.tolist(), [3.5, 20.0]) dist_mock.all_reduce.assert_called_once() self.assertIs(dist_mock.all_reduce.call_args.kwargs["op"], dist.ReduceOp.AVG) # The reduced buffer is the returned tensor itself (in place).