Repository navigation
Conversation
Add deepspeed.sequence.vocab_parallel_ops with vocab_parallel_logprobs (single id or top-k), vocab_parallel_entropy and vocab_parallel_kl_div (forward/reverse). They operate on the local logits shard of a vocabulary-parallel LM head, so RL and on-policy distillation losses never all-gather the full vocabulary. Each op is a custom autograd Function with closed-form local gradients, two all-reduces of per-token statistics, and row-chunked forward and backward that keeps peak memory a small multiple of the shard. Signed-off-by: Jin, Youzhi <youzhi.jin@intel.com> Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Keep the global max and the log-sum-exp separate instead of folding them into one logsumexp. Adding a small normalizer to a large max rounds it away in fp32, so e.g. logits [1e8, 1e8] produced log-probabilities of 0 instead of -log 2. Treat zero-probability tokens as contributing nothing to entropy, KL and their gradients. Partially masking the vocabulary to -inf previously produced 0 * -inf = NaN, while a target putting mass on a token masked in the other distribution still yields an infinite KL. Compare bf16 results against the fp32 reference instead of only checking dtypes. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi <youzhi.jin@intel.com>
Decide which tokens contribute nothing from the log-probability being -inf rather than the probability being 0. A finite log-probability such as -200 underflows to 0 after exp in fp32, so the previous check turned KL(teacher || student) into 0 when the teacher still had support on a token the student masked out. Such support mismatches now yield +inf. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi <youzhi.jin@intel.com>
A row whose requested ids are all ignore_index has a constant 0 output, but backward multiplied its softmax by 0. For a fully -inf-masked padding row the softmax is NaN, so NaN leaked into the logits gradient. Fill the gradient of such rows with 0 explicitly. Also document that index must be identical across TP ranks and that KL follows fp32 F.kl_div semantics for underflowing probabilities. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi <youzhi.jin@intel.com>
Codex Review SummaryThis comment shows the latest Codex review activity on this pull request.
ℹ️ About Codex in GitHubYour team has set up Codex to review pull requests in this repo. Reviews are triggered when you
Codex reacts with 👀 while any review is running, comments if it has suggestions, and reacts with 👍 once all reviews finish with no findings. |
There was a problem hiding this comment.
💡 Codex Review
Here are some automated review suggestions for this pull request.
Reviewed commit: ee94d8ee7e
ℹ️ About Codex in GitHub
Codex has been enabled to automatically review pull requests in this repo. Reviews are triggered when you
- Open a pull request for review
- Mark a draft as ready
- Comment "@codex review".
If Codex has suggestions, it will comment; otherwise it will react with 👍.
When you sign up for Codex through ChatGPT, Codex can also answer questions or update the PR, like "@codex address that feedback".
| @@ -0,0 +1,285 @@ | |||
| # SPDX-License-Identifier: Apache-2.0 | |||
There was a problem hiding this comment.
Add the required Signed-off-by trailer
This is a single-parent, non-merge commit, and git log -1 --format='%B' d68555b1a0a356e91fa8c9fd16b7557cb890c008 shows no Signed-off-by: trailer. The required DCO check will reject it, so recreate the commit with --signoff before merging.
AGENTS.md reference: AGENTS.md:L8-L9
Useful? React with 👍 / 👎.
There was a problem hiding this comment.
Thanks. This looks like a false positive: d68555b is not part of this PR. All commits on the branch (8d853fb, 99261be, f2ade50, ee94d8e, e31b42c) carry a Signed-off-by trailer, and the DCO check passes.
0z5a
left a comment
There was a problem hiding this comment.
Reviewed ee94d8e. I reproduced two correctness/error-handling gaps: output masking does not prevent NaN padding gradients in entropy and KL, and an empty shard makes TP ranks fail at different stages. I would address these before merging. I also have a follow-up suggestion for reusing validated shard metadata in repeated calls.
The probes execute the original function/autograd bodies on CPU, bridging DeepSpeed communicator calls to real Gloo at TP=1/2. Valid-row values and gradients match an independent full-vocabulary PyTorch oracle. These are not results from the complete DeepSpeed installation, XPU/NCCL backend, or CUDA Graph suite.
| sum_exp = torch.empty_like(global_max) | ||
| for i, logits in enumerate(logits_2d_list): | ||
| for rows in _row_chunks(logits): | ||
| sum_exp[i, rows] = torch.exp(logits[rows].float() - global_max[i, rows].unsqueeze(-1)).sum(dim=-1) |
There was a problem hiding this comment.
[P1] Prevent NaN gradients from fully masked padding rows. With one valid row and one all--inf row, I reproduced this for entropy and both KL directions at TP=1 and TP=2 on CPU/Gloo: after replacing the padding output with zero using masked_fill_, the output and scalar loss are finite, but backward still returns NaNs for every padding logit. Both student and teacher gradients are affected for KL. The valid row matches a full-vocabulary oracle, and the ignored-row logprobs control returns finite zero padding gradients.
Here global_max=-inf makes the shifted logits NaN; the saved statistics retain those NaNs when the caller masks the returned output, and zero upstream gradients do not remove them. Could we define and enforce a consistent ignored-row contract for entropy/KL and add masked forward/backward regressions? If fully masked rows are unsupported, please reject them safely before the unsafe distributed computation and establish that invariant at the calling boundary.
There was a problem hiding this comment.
Thanks for the repro. Fixed in e31b42c. A fully -inf row has no defined distribution, so entropy/KL (and logprobs for a valid index) still return NaN for it. That keeps an unmasked mistake visible. Its logits now get zero gradient in every op, including both KL inputs, so masking such rows (e.g. padding) is safe. I documented this contract in the module docstring. test_masked_undefined_row_gets_zero_gradient covers all four ops. It checks that the padded row gets zero gradient, the valid row matches its unpadded computation, and the teacher gradient stays finite for KL.
| if (vocab_start_index is None) != (vocab_end_index is None): | ||
| raise ValueError("vocab_start_index and vocab_end_index must be provided together") | ||
|
|
||
| vocab_start_index, _, global_vocab_size = _resolve_vocab_metadata(vocab_parallel_logits.shape[-1], |
There was a problem hiding this comment.
[P2, follow-up] Provide a path for reusing validated shard metadata. The PR documents per-call validation. In a TP=2 Gloo probe, three logprobs calls with unchanged explicit shard bounds still perform three metadata all-gathers, three .tolist() calls and three .item() calls, in addition to the softmax/statistics collectives. Inferred bounds add six metadata MIN/MAX all-reduces and increase .item() calls to nine. The existing trusted-metadata CE path avoids those metadata collectives.
Could repeated-call users reuse the instance-owned validation pattern from VocabParallelCausalLMLoss, with the group/device/shard lifetime bound to that owner? An optional validated-metadata path would keep the standalone checked API available without introducing a process-global cache.
For CUDA indices, the separate invalid_index.any().item() also prevents graph capture, so metadata caching alone would not establish capture support. Please keep input validation explicit when defining a steady-state path. This can be a follow-up if these wrappers intentionally remain eager-only; I have not measured GPU/NCCL latency or run CUDA capture.
There was a problem hiding this comment.
Added in e31b42c. vocab_parallel_logprobs now accepts global_vocab_size together with explicit shard bounds, mirroring vocab_parallel_cross_entropy. A repeated caller can validate the layout once (e.g. via _resolve_vocab_metadata on an object it owns) and skip the per-call metadata collective, without a process-global cache. The index range check still runs on every call, as you suggested. The .item() sync it implies remains, so these wrappers stay eager-only. CUDA graph capture is out of scope for this PR.
|
|
||
| def vocab_parallel_entropy(vocab_parallel_logits, tp_group=None): | ||
| """Per-token entropy of the softmax over vocabulary-sharded logits, in fp32.""" | ||
| entropy = _VocabParallelEntropy.apply(vocab_parallel_logits, tp_group) |
There was a problem hiding this comment.
[P2] Reject unsupported empty shards consistently across TP ranks. Entropy and KL enter the compute path without the collective shard validation used by logprobs/CE. With V_local=0 on rank 0 and 3 on rank 1, the original CPU/Gloo paths fail asymmetrically: rank 0 raises in _flatten at reshape(-1, 0) before issuing a collective, while rank 1 enters the MAX all-reduce and fails only when the probe's 3-second process-group timeout expires. The logprobs control instead gives both ranks the same clear empty-shard ValueError after metadata gathering.
Could we enforce the nonempty-shard invariant collectively before entropy/KL compute, reusing the existing validation contract or caller-owned validated metadata? A local-only assertion would retain the asymmetric failure. Please add a bounded TP=2 regression verifying coordinated rejection without leaving a peer waiting in a compute collective.
There was a problem hiding this comment.
Fixed in e31b42c, with a slightly different contract. Entropy and KL do not need shard bounds, so instead of adding a validation collective, an empty local shard now joins the same MAX/SUM all-reduces, contributing -inf and 0. The result is correct, no rank fails early, and no rank is left waiting. _flatten no longer uses reshape(-1, 0). logprobs keeps rejecting empty shards on all ranks through the existing metadata validation. TestVocabParallelOpsTP::test_empty_shard_does_not_desynchronize_ranks (TP=2, rank 0 owns no vocabulary) checks entropy/KL values and gradients against the full-vocabulary oracle, and checks that logprobs raises on both ranks. On the previous code, this test hangs.
…llel ops A fully -inf row (e.g. padding) has no softmax: its global max is -inf and the shifted logits are NaN, so masking its output still left NaN gradients for entropy and KL. Keep the NaN forward value but zero that row's gradient in every op, including both KL inputs. Entropy and KL need no shard bounds, so an empty local shard now joins the MAX/SUM collectives with -inf and 0 instead of failing in reshape on one rank while its peer waits in the collective. logprobs keeps rejecting empty shards on all ranks through the existing metadata validation. Let logprobs accept a trusted global_vocab_size with explicit shard bounds, matching vocab_parallel_cross_entropy, so repeated callers can skip the per-call metadata collective. Index range checks still run on every call. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Signed-off-by: Jin, Youzhi <youzhi.jin@intel.com>
Implements PR-1 of #8790.
What
This PR adds
deepspeed.sequence.vocab_parallel_ops. The ops compute per-token quantities on the local logits shard of a vocabulary-parallel LM head (#8173), without gathering the full vocabulary:vocab_parallel_logprobs(logits, index, ...): log-probs of one id per token, or of top-k ids (index[..., k]). Positions equal toignore_indexreturn 0.vocab_parallel_entropy(logits, ...)vocab_parallel_kl_div(student, teacher, reverse=False, ...):KL(teacher‖student), orKL(student‖teacher)withreverse=True. Gradients flow to both inputs.How
[N]-sized per-token statistics: the max, the sum-exp, and then the selected logits, entropy, or KL. The KL op packs the student and teacher statistics into the same collectives.vocab_parallel_logprobsalso validates the shard layout with a small metadata all-gather per call, plus a MIN/MAX pair when the bounds are not passed explicitly. A caller that already validated the layout can passglobal_vocab_sizewith explicit bounds to skip that collective, as withvocab_parallel_cross_entropy.-infcontribute nothing to entropy, KL, or their gradients. A fully-infrow (e.g. padding) has no distribution: it returns NaN but gets zero gradient, so masking its output is safe. Entropy and KL need no shard bounds and also work when a rank's local shard is empty.[N]statistics are saved.Tests
tests/unit/v1/sequence_parallelism/test_vocab_parallel_ops.pyhas 33 tests. Values and gradients are checked against a full-vocab PyTorch oracle in fp32. bf16 inputs are checked against the fp32 oracle with a tolerance.-inf-masked vocabularies (oracle: the reduced vocabulary), and infinite KL when the target has mass on a masked token, including when that mass underflows to 0 in fp32, and zero (not NaN) gradients for ignored or masked fully-infrows in all four ops;global_vocab_sizepath, and double backward raising an error;Ran on 8× Intel Arc Pro B70 (XPU, XCCL): this file plus
test_vocab_parallel_cross_entropy.py, 51 passed. TP=2 tests cover uneven shards and an empty shard on one rank.Memory
1× B70, T=8192, V_local=16032 (128K / TP8), bf16. fwd+bwd peak, excluding inputs:
Limitations