Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 2 additions & 9 deletions tzrec/loss/listwise_rank_loss.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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.

Expand All @@ -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
Expand Down Expand Up @@ -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)
12 changes: 0 additions & 12 deletions tzrec/loss/listwise_rank_loss_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
27 changes: 20 additions & 7 deletions tzrec/models/dlrm_hstu.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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,
Expand Down
67 changes: 67 additions & 0 deletions tzrec/models/dlrm_hstu_onerank_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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,
Expand Down Expand Up @@ -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:

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Minor / optional — two cheap hardenings for the new gate:

  • The mock only exercises the flag-on model. If the global_average guard in DlrmHSTU.loss were dropped, every single-process factor still collapses to 1.0, so nothing in the suite would catch that regression. Recomputing base.loss inside the same mock and asserting it equals base_losses (or asserting dist_mock.all_reduce.call_count == 4) would pin the gate.
  • _has_listwise_loss is an any() over all tasks' losses, but the listwise loss is wired only into is_click (as in every listwise test in this file). A variant with the listwise term on a later task — or on two tasks, which would also exercise the single request_avg_weight being reused across multiple listwise terms — would cover the actual scan.

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.
Expand Down
22 changes: 4 additions & 18 deletions tzrec/models/rank_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Minor — the contract behind the deleted guard is now undocumented.

The removed loss_weight = None carried a comment explaining why ("it is sized off the candidate count, not the request count"). After this change the shared tail applies whatever loss_weight the caller passes to the request-level scalar. Nothing can go wrong today — only DlrmHSTU.loss reaches this branch (only the HSTU family publishes TARGET_REPEAT_INTERLEAVE_KEY) and it passes only None or the request-count scalar — but a future model that publishes the key and reuses RankModel.loss/MultiTaskRank.loss with per-sample weights would silently broadcast a candidate-sized vector against the request-level scalar instead of failing. A one-line # NOTE: at this branch preserving the deleted rationale would keep the invariant stated where it's enforced, at no structural cost.

else:
raise ValueError(f"loss[{loss_type}] is not supported yet.")
if loss_weight is not None:
Expand Down
27 changes: 21 additions & 6 deletions tzrec/utils/fx_util.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
57 changes: 29 additions & 28 deletions tzrec/utils/fx_util_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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``.
"""

Expand All @@ -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,
Expand Down Expand Up @@ -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).
Expand Down
Loading