Skip to content

feat: shared-prefix GRPO training - #4008

Draft
zaifengp-maker wants to merge 11 commits into
NVIDIA-NeMo:mainfrom
zaifengp-maker:zaifengp/prototype-grpo-dedup
Draft

zaifengp-maker wants to merge 11 commits into
NVIDIA-NeMo:mainfrom
zaifengp-maker:zaifengp/prototype-grpo-dedup

Conversation

@zaifengp-maker

Copy link
Copy Markdown

What does this PR do ?

Forwards a GRPO group's shared prompt once instead of once per rollout.

Every rollout in a GRPO group shares the same prompt, so today the prompt's
forward pass is recomputed for each completion in the group. This
deduplicates it: the shared prefix is forwarded once and reused across the
group's completions.

  • nemo_rl/models/automodel/shared_prefix.py — the shared-prefix forward path.
  • nemo_rl/models/automodel/shared_prefix_moe.py — Qwen3 MoE expert-parallel
    variant, including logical routing bias.
  • nemo_rl/data/packing/algorithms.py — shared-prefix-aware sequence packing.
  • Chunked shared-prefix target logprobs, TP support, and the plumbing through
    automodel/setup.py, train.py, batched_data_dict.py, model_utils.py,
    lm_policy.py and dtensor_policy_worker_v2.py.
  • ~5000 lines of unit tests, including distributed EP and TP cases.

Issues

None filed yet. Happy to open one to hold the design discussion if that is
preferred over discussing it here.

Usage

Off by default; training behavior is unchanged unless it is enabled. Requires
DTensor v2 and sequence packing.

policy:
  # Dense Qwen3/Llama use force_hf=true; native Qwen3-MoE EP uses force_hf=false.
  shared_prefix_training: True

Bin packing charges each prompt once per group/bin; logprob inference stays
dense. See examples/configs/grpo_math_1B.yaml in this PR.

Before your PR is "Ready for review"

Pre checks:

  • Make sure you read and followed Contributor guidelines
  • Did you write any new necessary tests? — ~5000 lines under tests/unit/models/automodel/ and tests/unit/models/policy/, including distributed EP/TP cases
  • Did you run the unit tests and functional tests locally? — not against current main. The distributed tests need multi-GPU.
  • Did you add or update any necessary documentation? — no. Not written yet, pending direction on whether this belongs in tree.

Zaifeng Pan and others added 11 commits August 10, 2026 11:24
GRPO samples N responses per prompt, so a training microbatch repeats each
prompt N times. Train the DTensor v2 / HF / FA2 path on a compact
[1, compact_tokens] representation instead: one copy of each prompt followed
by every response that shares it. Qwen3 norm, MLP, and projections run
directly on that representation, and a registered attention backend preserves
the logical causal semantics with two varlen FA2 calls (prompt self-attention,
and response attending to prompt + its own response). Logprob inference stays
dense; only the train step is compacted.

The logical [batch, sequence] batch is unchanged. grpo.py infers prompt
lengths from the token loss mask, crops the terminal environment observation,
and attaches per-row group metadata; the worker builds the compact layout,
runs the forward with logits_to_keep restricted to predictor positions, and
scatters response logprobs back into logical order before the unmodified loss
function consumes them.

Support is deliberately narrow and validated loudly at setup: Qwen3 text-only,
DTensor v2, sequence packing on, force_hf, TP=CP=1, no dynamic batching, and
logprob-style losses without top-k/top-p or fused linear logprobs.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Zaifeng Pan <zaifengp@oci-hsg-cs-001-login-01.cm.cluster>
Replace the single full-width log_softmax over the compact response logits
with a custom autograd Function that computes target-only logprobs one chunk
at a time. Forward saves just the (model-owned) logits and the target ids and
releases each fp32 chunk immediately; backward rematerializes one fp32 softmax
chunk at a time instead of retaining response_tokens * vocab_size fp32
activations across the model backward.

Values and gradients are unchanged — the chunk size is a pure memory/compute
trade-off driven by policy.logprob_chunk_size, and defaults to a single chunk
when unset.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Zaifeng Pan <zaifengp@oci-hsg-cs-001-login-01.cm.cluster>
Bin packing charged every rollout its full dense length, so a shared-prefix
microbatch paid for N copies of a prompt it only stores once. Add a packer
that charges a prompt once per group per bin: rows of a group are first split
into the minimum number of first-fit response fragments, and those fragments
are then packed without splitting so another group's tail gap cannot force an
avoidable prompt recomputation. A group that still spans bins re-pays its
prompt in each, which the compact layout already handles.

shard_by_batch_size routes to the new packer when the caller names the group-id
and prompt-length columns, and measures bin length with the same deduplicated
cost. It also packs on real input lengths rather than TP-aligned ones, since
compact execution drops that alignment tail and charging it would turn padding
into fake response tokens.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Zaifeng Pan <zaifengp@oci-hsg-cs-001-login-01.cm.cluster>
_add_shared_prefix_training_metadata cropped input_lengths to the last
trainable token, dropping the terminal environment observation that native
rollouts append. The crop is numerically safe — that tail has a zero loss mask
and sits after every trainable token, so nothing depends on it — and it does
shrink the compact microbatch.

But it is an optimization orthogonal to prefix deduplication: pack_sequences
keeps the same tail on the dense path, so leaving the crop enabled gives
shared-prefix a packing advantage that a matched A/B would miscredit to
deduplication. Disable it so both paths consume identical input_lengths, and
re-enable once the crop is applied to both paths or measured separately.

The assignment is commented out rather than removed so the restoration point
stays obvious; the unit test now asserts the uncropped lengths.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Signed-off-by: Zaifeng Pan <zaifengp@oci-hsg-cs-001-login-01.cm.cluster>
Signed-off-by: Zaifeng Pan <zaifengp@oci-hsg-cs-001-login-01.cm.cluster>
Signed-off-by: Zaifeng Pan <zaifengp@oci-hsg-cs-001-login-01.cm.cluster>
Signed-off-by: Zaifeng Pan <zaifengp@oci-hsg-cs-001-login-01.cm.cluster>
Signed-off-by: Zaifeng Pan <zaifengp@oci-hsg-cs-001-login-01.cm.cluster>
Signed-off-by: Zaifeng Pan <zaifengp@oci-hsg-cs-001-login-01.cm.cluster>
Signed-off-by: Zaifeng Pan <zaifengp@oci-hsg-cs-001-login-01.cm.cluster>
Signed-off-by: Zaifeng Pan <zaifengp@oci-hsg-cs-001-login-01.cm.cluster>
@copy-pr-bot

copy-pr-bot Bot commented Sep 4, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

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

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant