Repository navigation
[BUG] GRPOLoss masking strategies produce incompatible shapes #4227
Description
Activity
Hi, I picked this up and reproduced the whole matrix on CPU with the tiny cached checkpoint
trl-internal-testing/tiny-Qwen2ForCausalLM-2.5plus theQwen/Qwen2.5-0.5Btokenizer, no vLLM and no generation (response tokens are fabricated, which is enough to characterize the loss side). Script attached below. Results:Strategy Input mode Outcome Mask used sft tokens finite loss response mask from prompt positions generic tokens finite loss attention mask rlhf tokens ValueError assistant mask not computable sft history finite loss falls back to assistant mask rlhf history finite loss assistant mask from chat template generic history finite loss attention mask Three observations:
-
Where the supported combinations run, the shape contract is actually consistent: masks are
(B, T)booleans over the full padded sequence, the distribution log prob is(B, T), and a per sequence advantage(B, 1, 1)broadcasts fine. I could not reproduce a shape bug in the loss reduction itself. -
I believe the real blocker for the skipped test is the transformers wrapper history path at batch greater than one, not a mismatch between strategies. With a prompt of 2 messages and a response of 1 message: batch 1 tokenizes cleanly and produces a correct assistant mask; batch 2 raises
RuntimeError: stacking tensordicts requires them to have congruent batch sizes, got td[1].batch_size=torch.Size([1]) and td[0].batch_size=torch.Size([2]); with congruent message counts it gets past that and fails later inside transformersmasking_utilswith a dimension mismatch. The skipped test uses batch 2 with heterogeneous prompts, which matches the first failure. Happy to file that as a separate issue with a probe script, since it also affects SFT with multi message histories. -
Two contract questions where the current behavior is undocumented:
rlhfwith tokens input can never work (the wrapper does not compute("masks", "all_assistant_mask")in that mode) and dies with the generic "Assistant mask not found" error. Should the contract be: rlhf requires history input or a caller supplied assistant mask, with an error that says exactly that?sftwith history input silently falls back to the assistant mask (_get_dist_with_prompt_maskinpolicies/common.py), which makessftandrlhfproduce identical distributions. Should that fall back warn, or is it intended?
Proposed plan, following your framing: first state the contract in the
GRPOLossdocstring (the matrix above plus shapes), then add a small deterministic CPU test covering the matrix so any future shape divergence fails in CI instead of in a skipped integration test. I would rather not touch the loss math, since the evidence says the math is fine.Repro: the two scripts are in my fork under
repro_4227.py(strategy by input mode matrix) andprobe_history.py(batch 1 versus batch 2 wrapper probe). Both run on CPU in about a minute withHF_HUB_OFFLINE=1once the two tiny artifacts are cached.Would you confirm the intended semantics for the two questions above? Then I will send the docstring plus deterministic test as one small PR referencing this issue.
-
Thanks for sharing your findings. It seems like the shape issue might not be with the loss reduction itself, but rather with the transformers wrapper when handling history mode at batch sizes greater than one. This is helpful information. Do you have any suggestions on how we might address the batch size issue with the transformers wrapper?
Thanks for the reproduction and for #4340.
I agree the original “incompatible loss shapes” report does not hold. The remaining work on this issue is the documented strategy contract plus the deterministic tests in that PR.
The TransformersWrapper history path at batch>1 looks like a separate bug (stacking / masking_utils). I would keep that out of #4227 and follow whatever issue you open for it, rather than expanding the scope here.
Describe the bug
GRPOLossis documented to supportmasking_strategyin{"sft", "rlhf", "generic"}, but the test that actually compares those strategies is skipped because of a shape mismatch. CI therefore will not catch a regression in the public loss when the mask strategy changes.This issue asks for the intended output shape per strategy and for the test to be unskipped once the loss matches that contract. It is not a request to land a speculative fix.
To Reproduce
The skip sits on
TestGRPOLoss.test_grpo_loss_with_real_models:The test builds tokens-mode input for
"sft"and history-mode input for"rlhf"(lines 1413-1442), generates withvLLMWrapper(..., return_log_probs=True), then constructsand expects a finite
result.loss_objective.On the loss side,
masking_strategyselects different distribution helpers:The class docstring already warns that a mismatch with the advantage mask produces shape errors (
grpo.py:377-383):Because the only cross-strategy test is skipped, that warning is not enforced.
Expected behavior
Please state the intended output shape of
GRPOLoss(and of the intermediate log-prob / mask tensors) for each of"sft","rlhf", and"generic", including how it should interact with tokens vs history input and with the advantage tensor.Once the loss matches that contract, unskip
test_grpo_loss_with_real_models(or replace it with a smaller deterministic test that still fails if the shapes diverge again).Screenshots
N/A
System info
Observed on a source checkout of pytorch/rl at
1d3de3db. The skipped test also needsvllmandtransformers.Additional context
I am not proposing a particular reduction or broadcast rule here. The skip reason is the whole report: the public API claims three strategies, and one of them currently cannot be compared to the other in CI.
Reason and Possible fixes
Unskip only after the intended shapes are specified and the loss (or the test inputs / advantage) is aligned to them. A drive-by squeeze / unsqueeze without that spec is likely to hide the next mismatch.
Checklist