fix(loss): preserve gradients for skipped listwise reranker batches - #10090
Merged
tastelikefeet merged 1 commit intoSep 15, 2026
Merged
tastelikefeet merged 1 commit into
tastelikefeet merged 1 commit into
Conversation
tastelikefeet
approved these changes
Sep 15, 2026
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Listwise reranker loss returns a new, unrelated zero tensor when no query group is retained. With the documented LISTWISE_RERANKER_MIN_GROUP_SIZE filter, one DDP rank can skip its groups while another still contributes a loss. The skipped rank produces no model gradients, leaving its peer waiting in backward.
Return a sum over an empty logits slice in both zero-loss branches. This preserves the autograd connection and supplies zero gradients without reading skipped scores. Retained query losses and averaging are unchanged.
Validation: four CPU tests pass, covering both zero branches, three floating dtypes, skipped nonfinite scores, retained loss/gradient equivalence and real two-process DDP with skipped and unequal-rank steps plus no_sync gradient accumulation. A separate time-bounded Gloo probe reproduces the upstream backward timeout and verifies identical gradients after the fix. Applicable pre-commit checks pass. No full pretrained-model or Megatron/NCCL training run was performed.
Additional verification uses 600 unchanged public ranking records from SciDocs, StackOverflow and AskUbuntu through the existing preprocessor and reranker grouping/collation methods, with a tiny deterministic text-feature scorer. Across three minimum-size settings, 1,762 retained cases match upstream loss and gradients exactly; 38 skipped cases now supply zero gradients. Original AskUbuntu records reproduce the asymmetric-rank backward timeout at the default threshold of 2; the fix completes successfully. These are grouping/autograd checks, not pretrained-model quality tests.
When an entire optimizer step has zero gradients, existing optimizer momentum or weight decay can still update parameters. This change restores connected-zero semantics; it does not add globally empty-step skipping. AdamW and momentum SGD match an explicit connected-zero reference.
A random one-layer BERT smoke test also reproduces the upstream timeout with find_unused_parameters both enabled and disabled. The fix passes both configurations and non-reentrant gradient checkpointing; all 25 parameter gradients match a cross-entropy reference. This is a CPU/Gloo model-structure check, not full Trainer.train or pretrained training.