-
Notifications
You must be signed in to change notification settings - Fork 85
[bugfix] apply enable_global_average_loss to the list-wise rank loss #678
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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) | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Minor — the contract behind the deleted guard is now undocumented. The removed |
||
| else: | ||
| raise ValueError(f"loss[{loss_type}] is not supported yet.") | ||
| if loss_weight is not None: | ||
|
|
||
There was a problem hiding this comment.
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:
global_averageguard inDlrmHSTU.losswere dropped, every single-process factor still collapses to 1.0, so nothing in the suite would catch that regression. Recomputingbase.lossinside the same mock and asserting it equalsbase_losses(or assertingdist_mock.all_reduce.call_count == 4) would pin the gate._has_listwise_lossis anany()over all tasks' losses, but the listwise loss is wired only intois_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 singlerequest_avg_weightbeing reused across multiple listwise terms — would cover the actual scan.