Conversation
|
Adding a runnable counter-example, and a note that narrows the blast radius stated in the description. Counter-exampleThe batch below carries an identifiable marker per row, so a misalignment can be read straight off the numbers. Rewards are encoded as Row 2 is the clearest case. The sample is Note Narrowing the blast radiusI traced every consumer of
So the corrupted 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 The corroboration in the description still stands and is the strongest argument here: |
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>
990739a to
9f9bb3f
Compare
Problem
dynamic_sampling()filters a generation batch down to the samples with non-zero reward std (DAPO). Right after filtering, it does this:select_indices()already produced a correctly filtered, row-alignedtotal_rewardonfiltered_repeated_batch. The code above immediately overwrites it withtotal_rewards, the full, unfiltered per-round reward tensor captured before filtering. Every other key onfiltered_repeated_batch(message_log,std,baseline, and the just-renamedfiltered_reward) reflects only the kept samples.total_rewardand the rest of the batch only happen to agree when a round keeps every generated sample; in every other case they diverge, silently:train_prompts_sizeneeds, the code doesfiltered_repeated_batch.slice(0, train_prompts_size).BatchedDataDict.slice()slices each key independently against its own tensor, so this takes the firsttrain_prompts_sizerows oftotal_reward's full, unfiltered array (which still contains the dropped, zero-std samples) instead of the firsttrain_prompts_sizekept rows the rest of the batch reflects.total_rewardends up pairing the wrong reward with each kept sample; it can even include a reward that belongs to a sample DAPO decided to drop.BatchedDataDict.from_batches([batch_cache, filtered_repeated_batch]).from_batchesconcatenates each key independently too:total_rewardaccumulates using each round's full per-round length while every other key accumulates using the kept-only length, sototal_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'sselect_indices/slice/from_batchesall 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 corruptedtotal_reward:metrics["reward"](the DAPO health metric) and the per-step training-data JSONL (log_data["rewards"]), both underif 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:A rename, not an overwrite, so
total_rewardstays 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 theBatchedDataDict, rather than smuggling it into the same dict under a mismatched row count. This PR aligns the legacydynamic_sampling()path with that same, already-correct pattern.Fix
In
dynamic_sampling(): stop overwritingtotal_reward. Rename it tofiltered_rewardinstead (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 formetrics["reward"], which is a deliberately distinct diagnostic frommetrics["filtered_reward"](raw generation quality vs. reward of the samples that end up training). Sincerepeated_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, beforedynamic_sampling()reassignsrepeated_batch, and use that name formetrics["reward"]instead.log_data["rewards"](the per-step training-data JSONL, under the sameuse_dynamic_samplingguard) keeps readingrepeated_batch["total_reward"], unchanged. It is logged row-for-row againstcontent/token_ids, which reflect the filtered-and-sliced batch, so it needs the filtered reward, unlikemetrics["reward"]. Since this now makeslog_data["rewards"]identical tolog_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 usingthis_round_unfiltered_rewardshere would raise insidelog_batched_dict_as_jsonlwhenever its length (one generation batch) differs from the accumulated-and-slicedcontent's length, which is the common case under multi-round dynamic sampling.Evidence
Three tests, isolating each divergent case and the caller-side fix:
test_dapo_dynamic_sampling_discard_slice_preserves_reward_alignment: a single round generates 6 samples, one dropped (zero std), 4 needed. Asserts the returnedtotal_rewardequalsfiltered_rewardand both equal the first 4 of the 5 kept rewards, in order.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 finaltotal_rewardequalsfiltered_rewardand both equal the correct 6 accumulated-and-sliced kept rewards.test_grpo_train_reward_metric_uses_preserved_unfiltered_rewards: locksgrpo_train'smetrics["reward"]assignment to the preserved pre-filter variable via source inspection (grpo_trainbuilds 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:With the fix applied, all three pass:
The full
test_grpo.pysuite (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:Scope
This only touches
dynamic_sampling()'s handling oftotal_reward/filtered_rewardand the one place ingrpo_trainthat reads the pre-filter reward for logging. Everything else about how dynamic sampling selects, caches, and slices prompt groups is unchanged.