Skip to content

fix(grpo): stop corrupting total_reward across dynamic-sampling rounds - #3794

Open
khazic wants to merge 3 commits into
NVIDIA-NeMo:mainfrom
khazic:khazic/fix/dynamic-sampling-reward-misalignment
Open

khazic wants to merge 3 commits into
NVIDIA-NeMo:mainfrom
khazic:khazic/fix/dynamic-sampling-reward-misalignment

Conversation

@khazic

@khazic khazic commented Aug 24, 2026

Copy link
Copy Markdown

Problem

dynamic_sampling() filters a generation batch down to the samples with non-zero reward std (DAPO). Right after filtering, it does this:

filtered_repeated_batch = repeated_batch.select_indices(keep_prompt_indices)
filtered_repeated_batch["std"] = std[keep_prompt_indices]
filtered_repeated_batch["baseline"] = baseline[keep_prompt_indices]

# Store filtered and total rewards to track them separately
filtered_rewards = filtered_repeated_batch["total_reward"]
filtered_repeated_batch["total_reward"] = total_rewards
filtered_repeated_batch["filtered_reward"] = filtered_rewards

select_indices() already produced a correctly filtered, row-aligned total_reward on filtered_repeated_batch. The code above immediately overwrites it with total_rewards, the full, unfiltered per-round reward tensor captured before filtering. Every other key on filtered_repeated_batch (message_log, std, baseline, and the just-renamed filtered_reward) reflects only the kept samples. total_reward and the rest of the batch only happen to agree when a round keeps every generated sample; in every other case they diverge, silently:

  • Single round, over-generation. When more samples survive filtering than train_prompts_size needs, the code does filtered_repeated_batch.slice(0, train_prompts_size). BatchedDataDict.slice() slices each key independently against its own tensor, so this takes the first train_prompts_size rows of total_reward's full, unfiltered array (which still contains the dropped, zero-std samples) instead of the first train_prompts_size kept rows the rest of the batch reflects. total_reward ends up pairing the wrong reward with each kept sample; it can even include a reward that belongs to a sample DAPO decided to drop.
  • Multiple rounds. When the buffer needs several generation batches to fill, the partial batches are merged via BatchedDataDict.from_batches([batch_cache, filtered_repeated_batch]). from_batches concatenates each key independently too: total_reward accumulates using each round's full per-round length while every other key accumulates using the kept-only length, so total_reward's row count actively diverges from the rest of the batch mid-accumulation, before the final slice forces it back to a coincidentally-matching row count with the wrong values inside it.

Neither case raises an error. BatchedDataDict's select_indices/slice/from_batches all operate per-key, so a wrong-sized or wrong-order tensor is silently sliced or concatenated down to a row count that matches everyone else's by the time the batch is returned; only the values are wrong. Two things downstream read this corrupted total_reward: metrics["reward"] (the DAPO health metric) and the per-step training-data JSONL (log_data["rewards"]), both under if master_config.grpo.use_dynamic_sampling.

Independent corroboration

grpo_sync.py's own reimplementation of this same algorithm for the data-plane (TQ) path, _apply_dynamic_sampling, already avoids this exact mistake:

survivors_carry["filtered_reward"] = survivors_carry["total_reward"]

A rename, not an overwrite, so total_reward stays the correctly filtered value there. It also tracks the unfiltered per-round reward for logging as a plain Python list (pending_unfiltered_rewards) kept entirely outside the BatchedDataDict, rather than smuggling it into the same dict under a mismatched row count. This PR aligns the legacy dynamic_sampling() path with that same, already-correct pattern.

Fix

In dynamic_sampling(): stop overwriting total_reward. Rename it to filtered_reward instead (filtered_repeated_batch["filtered_reward"] = filtered_repeated_batch["total_reward"]), so both keys hold the same, correctly filtered and row-aligned value going forward. total_rewards, the now-unused local variable holding the full per-round reward, is removed from the function.

At the one caller (grpo_train's sync loop): the pre-filter, per-round reward is still needed for metrics["reward"], which is a deliberately distinct diagnostic from metrics["filtered_reward"] (raw generation quality vs. reward of the samples that end up training). Since repeated_batch["total_reward"] no longer carries that value after the call, capture it under its own name (this_round_unfiltered_rewards) right where it is first read, before dynamic_sampling() reassigns repeated_batch, and use that name for metrics["reward"] instead.

log_data["rewards"] (the per-step training-data JSONL, under the same use_dynamic_sampling guard) keeps reading repeated_batch["total_reward"], unchanged. It is logged row-for-row against content/token_ids, which reflect the filtered-and-sliced batch, so it needs the filtered reward, unlike metrics["reward"]. Since this now makes log_data["rewards"] identical to log_data["filtered_rewards"] where they previously (incorrectly) differed, review flagged it as a suspicious-looking regression; it is documented at the call site and locked with a source-inspection test instead, because using this_round_unfiltered_rewards here would raise inside log_batched_dict_as_jsonl whenever its length (one generation batch) differs from the accumulated-and-sliced content's length, which is the common case under multi-round dynamic sampling.

Evidence

Three tests, isolating each divergent case and the caller-side fix:

  1. test_dapo_dynamic_sampling_discard_slice_preserves_reward_alignment: a single round generates 6 samples, one dropped (zero std), 4 needed. Asserts the returned total_reward equals filtered_reward and both equal the first 4 of the 5 kept rewards, in order.
  2. test_dapo_dynamic_sampling_cache_preserves_reward_alignment: two rounds are needed to fill a 6-sample buffer, round one drops one of its four samples. Asserts the final total_reward equals filtered_reward and both equal the correct 6 accumulated-and-sliced kept rewards.
  3. test_grpo_train_reward_metric_uses_preserved_unfiltered_rewards: locks grpo_train's metrics["reward"] assignment to the preserved pre-filter variable via source inspection (grpo_train builds this inline rather than through an extracted, directly callable helper, so invoking the real function would require standing up rollout, policy, and generation actors).

On the current main, all three fail:

FAILED tests/unit/algorithms/test_grpo.py::test_dapo_dynamic_sampling_discard_slice_preserves_reward_alignment
FAILED tests/unit/algorithms/test_grpo.py::test_dapo_dynamic_sampling_cache_preserves_reward_alignment
FAILED tests/unit/algorithms/test_grpo.py::test_grpo_train_reward_metric_uses_preserved_unfiltered_rewards

AssertionError: Tensor-likes are not close!
Mismatched elements: 4 / 6 (66.7%)
Greatest absolute difference: 10.0 at index (2,) (up to 1e-05 allowed)
...
AssertionError: grpo_train no longer sources metrics['reward'] from the reward captured before dynamic_sampling filtering; it would silently become identical to metrics['filtered_reward']

3 failed

With the fix applied, all three pass:

3 passed

The full test_grpo.py suite (141 pre-existing tests plus the 3 above, 144 total) passes on the fix branch, so the change does not regress any other GRPO path, including the other dynamic-sampling tests already covering the non-discarding case:

144 passed

Scope

This only touches dynamic_sampling()'s handling of total_reward/filtered_reward and the one place in grpo_train that reads the pre-filter reward for logging. Everything else about how dynamic sampling selects, caches, and slices prompt groups is unchanged.

@khazic
khazic requested review from a team as code owners August 24, 2026 06:19
@copy-pr-bot

copy-pr-bot Bot commented Aug 24, 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.

@khazic

khazic commented Aug 26, 2026

Copy link
Copy Markdown
Author

Adding a runnable counter-example, and a note that narrows the blast radius stated in the description.

Counter-example

The batch below carries an identifiable marker per row, so a misalignment can be read straight off the numbers. Rewards are encoded as prompt * 10 + generation. Prompt 1 has zero reward std, so DAPO drops rows 2 and 3. The three lines from dynamic_sampling() are reproduced verbatim.

input, one row per generated sample
  row 0: message_log=p0g0  total_reward=10  std=1
  row 1: message_log=p0g1  total_reward=11  std=1
  row 2: message_log=p1g0  total_reward=20  std=0
  row 3: message_log=p1g1  total_reward=21  std=0
  row 4: message_log=p2g0  total_reward=30  std=2
  row 5: message_log=p2g1  total_reward=31  std=2

rows kept by the non-zero-std filter: [0, 1, 4, 5]

after select_indices, before the overwrite: aligned
  row 0: message_log=p0g0  total_reward=10
  row 1: message_log=p0g1  total_reward=11
  row 2: message_log=p2g0  total_reward=30
  row 3: message_log=p2g1  total_reward=31

after the overwrite
  batch size (rows): 4
  len(total_reward): 6          <- longer than the batch it belongs to
  OK  row 0: message_log=p0g0  total_reward=10  (correct: 10)
  OK  row 1: message_log=p0g1  total_reward=11  (correct: 11)
  BAD row 2: message_log=p2g0  total_reward=20  (correct: 30)
  BAD row 3: message_log=p2g1  total_reward=21  (correct: 31)

Row 2 is the clearest case. The sample is p2g0, which earned 30, but total_reward reports 20, the reward of p1g0, a sample this round filtered out. With the PR applied all four rows stay aligned.

Note len(total_reward) is 6 while the batch has 4 rows. The column is not merely permuted, it belongs to a different row space, and it is only the later slice() that forces the count back into agreement while leaving the wrong values in place.

Narrowing the blast radius

I traced every consumer of total_reward after the overwrite, and the two paths that feed training numerics both read a correct value:

  • scale_rewards() is called before dynamic_sampling(), so it operates on the pre-overwrite tensor.
  • The advantage computation selects filtered_reward when use_dynamic_sampling is set, which is the correctly filtered tensor the same three lines saved under a new name.

So the corrupted total_reward reaches metrics["reward"], the per-step training-data JSONL, and the run's printed average reward, exactly as the description says. It does not reach advantages or reward scaling. This is a reporting-correctness bug, and I would rather state that precisely than leave the impression that advantages are affected.

That does not make it harmless. The DAPO health metric is what a user watches to decide whether dynamic sampling is behaving, and it currently reports rewards belonging to a different sample set, including samples the round discarded. It is also a trap for future code: any new consumer of total_reward on this path silently inherits the misalignment, with no error raised, because BatchedDataDict operates per key.

The corroboration in the description still stands and is the strongest argument here: grpo_sync.py's _apply_dynamic_sampling implements the same algorithm and already does a rename instead of an overwrite, keeping the unfiltered per-round rewards outside the batch entirely. This PR brings the older path in line with the one that is already correct.

@svcnvidia-nemo-ci svcnvidia-nemo-ci added the waiting-on-maintainers Waiting on maintainers to respond label Aug 28, 2026
dynamic_sampling() overwrote filtered_repeated_batch["total_reward"] with
total_rewards, the full, unfiltered per-round reward tensor, right after
select_indices() had already produced a correctly filtered, row-aligned
total_reward on the same object. Every other key in filtered_repeated_batch
(message_log, std, baseline, filtered_reward, ...) reflects only the kept
(non-zero std) samples, so the two only happen to agree when a single round
keeps every generated sample. In every other case they diverge:

- Single round, over-generation (more kept samples than train_prompts_size):
  the final .slice(0, train_prompts_size) takes an arbitrary prefix of the
  full, unfiltered reward array instead of the first train_prompts_size KEPT
  samples, so total_reward silently pairs the wrong reward with each kept
  sample, including rewards belonging to dropped, zero-std samples.
- Multiple rounds (buffer needs several generation batches to fill):
  BatchedDataDict.from_batches concatenates total_reward across rounds using
  each round's full per-round length while every other key concatenates using
  the kept-only length, so total_reward's row count actively diverges from the
  rest of the batch mid-accumulation.

Neither failure raises: BatchedDataDict.slice()/from_batches() operate
per-key, so the wrong-sized tensor is silently sliced/concatenated to a
row count that coincidentally matches everyone else's by the time the batch
is returned, while the values themselves stay wrong. metrics["reward"] and
the per-step training-data JSONL both read this corrupted total_reward.

grpo_sync.py's independent TQ-path reimplementation of this same algorithm
(_apply_dynamic_sampling) already gets this right: it renames the naturally
filtered value to filtered_reward instead of overwriting total_reward with
the unfiltered one, and tracks the unfiltered per-round reward separately
(pending_unfiltered_rewards) rather than smuggling it into the same
BatchedDataDict under a mismatched row count. Align dynamic_sampling() with
that pattern: total_reward becomes an alias of the correctly filtered value,
and the caller (grpo_train's sync loop) keeps the pre-filter, per-round reward
under its own name (this_round_unfiltered_rewards) for metrics["reward"],
so that diagnostic keeps showing the raw generation quality distinct from
metrics["filtered_reward"] instead of silently becoming identical to it.

Signed-off-by: khazic <khazzz1c@gmail.com>
Follow-up on review of the dynamic-sampling reward fix: log_data["rewards"]
under grpo.use_dynamic_sampling now reads the same, correctly filtered
total_reward as log_data["filtered_rewards"], since this PR fixed total_reward
to actually be the filtered, row-aligned value. That makes the two fields
identical where they previously (incorrectly) differed, which a reviewer
flagged as a suspicious-looking regression.

It is not one: log_data's other fields (content, token_ids, ...) all reflect
the filtered-and-sliced batch, so rewards must stay row-aligned to it.
this_round_unfiltered_rewards (used for metrics["reward"]) only covers the
single most recent generation batch, not the possibly multi-round accumulated
batch this log entry describes, so using it here would raise inside
log_batched_dict_as_jsonl whenever the two sizes differ. Document this
explicitly at the call site and lock both the metrics and log_data literals in
place with a source-inspection test, so a future "fix" restoring the
unfiltered value here does not reintroduce a crash.

Signed-off-by: khazic <khazzz1c@gmail.com>
…ing API

dynamic_sampling now requires is_trivial_prompt_distribution when
use_dynamic_sampling is set, and grpo_train delegates its training loop to
_grpo_train_impl. Pass the trivial-prompt mask the tests intend (std == 0)
and inspect _grpo_train_impl for the metrics and log_data literals.

Signed-off-by: khazic <khazzz1c@gmail.com>
@khazic
khazic force-pushed the khazic/fix/dynamic-sampling-reward-misalignment branch from 990739a to 9f9bb3f Compare September 24, 2026 15:16

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

community-request waiting-on-maintainers Waiting on maintainers to respond

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants