Skip to content

Add vocabulary-parallel logprobs, entropy and KL ops - #8791

Open
jinyouzhi wants to merge 6 commits into
deepspeedai:masterfrom
jinyouzhi:vocab-parallel-distill-ops
Open

jinyouzhi wants to merge 6 commits into
deepspeedai:masterfrom
jinyouzhi:vocab-parallel-distill-ops

Conversation

@jinyouzhi

@jinyouzhi jinyouzhi commented Oct 9, 2026 •

Copy link
Copy Markdown
Contributor

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 to ignore_index return 0.
  • vocab_parallel_entropy(logits, ...)
  • vocab_parallel_kl_div(student, teacher, reverse=False, ...): KL(teacher‖student), or KL(student‖teacher) with reverse=True. Gradients flow to both inputs.
# Policy-gradient OPD / GRPO
logp = vocab_parallel_logprobs(student_logits, sampled_ids, **ids)
# Teacher top-k forward KL (verl forward_kl_topk)
s_topk = vocab_parallel_logprobs(student_logits, teacher_topk_ids, **ids)
kl = (teacher_topk_logp.exp() * (teacher_topk_logp - s_topk)).sum(-1)
# Dense reverse KL (DeepSpeedExamples OPSD default)
kl = vocab_parallel_kl_div(student_logits, teacher_logits.detach(), reverse=True, tp_group=g)

How

  • Each op is a custom autograd Function. The forward does one MAX all-reduce and two SUM all-reduces of [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_logprobs also 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 pass global_vocab_size with explicit bounds to skip that collective, as with vocab_parallel_cross_entropy.
  • The global max and the log-sum-exp are kept separate, so a large common logit offset does not lose precision in fp32. Tokens whose log-probability is -inf contribute nothing to entropy, KL, or their gradients. A fully -inf row (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.
  • The backward is closed-form and local, with no communication. This follows the Megatron convention that every TP rank backprops the same loss.
  • Forward and backward process rows in chunks of at most 16M elements. Only the inputs and [N] statistics are saved.
  • Outputs are fp32 and replicated across the TP group.

Tests

tests/unit/v1/sequence_parallelism/test_vocab_parallel_ops.py has 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.

  • logprobs (single and top-k), entropy, forward/reverse KL, and teacher top-k forward KL;
  • forced chunking, empty input, in-place masking of outputs, detached teacher, bf16;
  • a large common logit offset (oracle: the unshifted logits), partially -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 -inf rows in all four ops;
  • argument validation, the trusted global_vocab_size path, and double backward raising an error;
  • TP=2 with uneven vocab shards.

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:

Op Naive PyTorch This PR Time (naive → this PR)
logprobs 1506 MiB 443 MiB 7.5 → 21.4 ms
entropy 3012 MiB 586 MiB 30.7 → 41.5 ms
reverse KL 3012 MiB 698 MiB 37.2 → 70.2 ms
forward KL 2510 MiB 634 MiB 24.8 → 64.3 ms

Limitations

  • Speed: up to ~2.9× slower than naive eager in these measurements, because softmax is recomputed in backward, work is chunked, and zero-probability tokens are masked explicitly. This PR trades time for memory and numerical robustness.
  • The local logits shard is still materialized. A fused-linear variant is future work.
  • Only first-order gradients are supported. Double backward raises an error.
  • Student and teacher must share the same TP group and vocab sharding.
  • No built-in reduction, mask, or temperature. These are applied by the caller.

jinyouzhi and others added 4 commits October 9, 2026 16:41
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>
@jinyouzhi
jinyouzhi marked this pull request as ready for review October 10, 2026 09:07
@chatgpt-codex-connector

chatgpt-codex-connector Bot commented Oct 10, 2026 •

Copy link
Copy Markdown

Codex Review Summary

This comment shows the latest Codex review activity on this pull request.

Review Status Commit Review trigger
📝 Code Review ✅ Completed 2026-10-10T09:14:52.199461Z ee94d8e Draft marked ready
ℹ️ About Codex in GitHub

Your team has set up Codex to 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" or "@codex security review".

Codex reacts with 👀 while any review is running, comments if it has suggestions, and reacts with 👍 once all reviews finish with no findings.

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

💡 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

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

P1 Badge 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 👍 / 👎.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

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 0z5a left a comment

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.

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)

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.

[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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

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],

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.

[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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

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)

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.

[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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

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.

jinyouzhi and others added 2 commits October 10, 2026 16:41
…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>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants