feat: shared-prefix GRPO training - #4008
Draft
zaifengp-maker wants to merge 11 commits into
Draft
zaifengp-maker wants to merge 11 commits into
zaifengp-maker wants to merge 11 commits into
Conversation
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>
This branch has not been deployed
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.
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-parallelvariant, including logical routing bias.
nemo_rl/data/packing/algorithms.py— shared-prefix-aware sequence packing.automodel/setup.py,train.py,batched_data_dict.py,model_utils.py,lm_policy.pyanddtensor_policy_worker_v2.py.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.
Bin packing charges each prompt once per group/bin; logprob inference stays
dense. See
examples/configs/grpo_math_1B.yamlin this PR.Before your PR is "Ready for review"
Pre checks:
tests/unit/models/automodel/andtests/unit/models/policy/, including distributed EP/TP casesmain. The distributed tests need multi-GPU.