Skip to content

fix(sampling): preserve per-row parameters in active chunks - #464

Open
MichaelCao0 wants to merge 3 commits into
RL-Align:mainfrom
MichaelCao0:codex/fix-sampling-per-row-params
Open

MichaelCao0 wants to merge 3 commits into
RL-Align:mainfrom
MichaelCao0:codex/fix-sampling-per-row-params

Conversation

@MichaelCao0

@MichaelCao0 MichaelCao0 commented Oct 3, 2026 •

Copy link
Copy Markdown
Contributor

vocab_parallel_sampling_keep_mask selects and chunks logits by original row IDs, but previously passed full-batch temperature, top_p, and top_k vectors into each chunk. Per-row parameters therefore failed with shape errors when the batch was split, compacted to active rows, or entirely inactive.

Normalize parameter shapes once and select per-row values using the same original row IDs as the logits. Scalars and singleton tensors retain broadcasting behavior; inactive rows remain unmasked and active vocabulary padding remains excluded. Document the contract and add independent per-row reference tests.

Validation on NVIDIA H200 with PyTorch 2.8.0+cu128:

  • Regression tests against the original implementation: 41 failed, 23 passed.
  • Fixed sampling tests: 64 passed; related sampling/rollout tests: 104 passed, including those 64.
  • Real NCCL TP=2 smoke: 180 cases per rank passed across FP32/BF16, scalar and per-row parameters, active subsets, empty selections, partial chunks, and padded vocabulary shards.
  • Changed-file Black, isort, Ruff, Flake8 and whitespace checks passed.

The helper currently has direct in-repository callers in the distributed validation script; this change does not claim a full-model training or vLLM integration result. Parameter indexing adds small per-chunk operations; no performance improvement is claimed.

Summary by CodeRabbit

  • Bug Fixes
    • Sampling now applies per-row parameters to the correct input rows when processing active chunks, while continuing to support shared scalar values.
    • Invalid parameter counts now raise a clear error instead of producing incorrect results.

Signed-off-by: MichaelCaoo <139663530+MichaelCao0@users.noreply.github.com>
@coderabbitai

coderabbitai Bot commented Oct 3, 2026 •

Copy link
Copy Markdown

Review in Change Stack →

Important

Review skipped

Review was skipped as selected files did not have any reviewable changes.

⚙️ Run configuration
  • Configuration used: defaults
  • Review profile: CHILL
  • Plan: Advanced
  • Run ID: 9b910033-77ac-4492-b52e-29c60e0da4e0
📥 Commits

Reviewing files that changed from the base of the PR and between e84d6ba and 0ad3a61.

You can disable this status message by setting the reviews.review_status to false in the CodeRabbit configuration file.

Use the checkbox below for a quick retry:

  • 🔍 Trigger review

No actionable comments were generated in the recent review. 🎉

ℹ️ Recent review info
⚙️ Run configuration
  • Configuration used: defaults
  • Review profile: CHILL
  • Plan: Advanced
  • Run ID: e1e296f4-2ce9-4029-8fb6-c5f6925fe721
📥 Commits

Reviewing files that changed from the base of the PR and between 6d8b4bc and e84d6ba.

📒 Files selected for processing (2)
  • rl_engine/integrations/sampling.py
  • tests/test_complete_sampling.py

Included review availability: This review used your included allowance. Your plan provides up to 2 included reviews per hour; 1 remain after this review.


📝 Walkthrough

Walkthrough

The sampling wrapper now accepts shared or per-original-row parameters, validates their sizes, and selects per-row values using original row indices for each chunk. Tests cover parameter forms, active-row selections, padding, invalid sizes, and empty batches.

Changes

Sampling parameter alignment

Layer / File(s) Summary
Parameter validation, selection, and coverage
rl_engine/integrations/sampling.py, tests/test_complete_sampling.py
The wrapper validates that each non-None parameter has one value or one value per input row. It selects per-row values by original row index and keeps singleton values shared. Tests cover parameter forms, chunk sizes, active rows, invalid sizes, and empty batches.

Priority: ⬇️ Low

Estimated code review effort: 2 (Simple) | ~10 minutes

Change: Bug fix

Merge Risk: ⚪ Minimal · up to e84d6

No actionable issue remains from this review; the change is mergeable after normal checks.

🚥 Pre-merge checks | ✅ 4 | ❌ 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 22.22% which is insufficient. The required threshold is 80.00%. Docstring coverage is scoped to functions touched by this diff. Analyzed 9 functions across 2 files. Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (4 passed)
Check name Status Explanation
Description Check ✅ Passed Check skipped - CodeRabbit’s high-level summary is enabled.
Title check ✅ Passed The title clearly describes the main change: preserving per-row sampling parameters when processing active chunks.
Linked Issues check ✅ Passed Check skipped because no linked issues were found for this pull request.
Out of Scope Changes check ✅ Passed Check skipped because no linked issues were found for this pull request.
✨ Finishing Touches 💡 1
🛠️ Fix failing CI checks 💡
  • Commit to this branch
  • Create a new PR
🧪 Generate unit tests (beta)
  • Create a new PR
  • Autopilot · Keep fixing CodeRabbit findings and required CI, and resolving merge conflicts

Comment @coderabbitai help to get the list of available commands.

@Flink-ddd Flink-ddd added the bug Something isn't working label Oct 4, 2026

@Flink-ddd Flink-ddd left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Could you share the commands, commit SHAs, and before and after test logs, including the TP=2 smoke test? A representative failure before the fix and its passing result afterward would help verify the change.

@MichaelCao0

Copy link
Copy Markdown
Contributor Author

I reran the before/after comparison against the current PR head. The attached pr-464-evidence-v2.zip contains the full logs, exact commands and exit codes, source hashes, and a standalone TP=2 reproduction script.

The revisions used were:

  • Baseline: 11cac8c46fa9f6fae67ffd3b4054c11337d63c04
  • Original fix: e84d6ba466de52a0eac0d1ed8de5727837535582
  • Current PR head: b20fbc5f957197dead2108a5a7d03d40a51c2b21

For the baseline run, I copied only tests/test_complete_sampling.py from the PR head onto the baseline source. Both runs therefore execute identical regression tests, while the baseline retains the original production implementation.

From each source directory, with PYTHONPATH="$PWD" and PYTEST_DISABLE_PLUGIN_AUTOLOAD=1:

python -m pytest tests/test_complete_sampling.py \
  -q --tb=short -p no:cacheprovider

# Fixed version: related tests, including the sampling file above.
python -m pytest tests/test_complete_sampling.py \
  tests/test_vllm_rollout_sampler.py \
  tests/test_vllm_worker_top_p_replay.py \
  -q --tb=short -p no:cacheprovider
Check Before After
Sampling test file 41 failed, 23 passed 64 passed
Related test set — 104 passed
Real NCCL TP=2 smoke Broadcast error on both ranks 180 cases passed per rank

A representative regression is test_chunked_per_row_parameters_follow_original_rows[temperature-all-1]. Seven per-row temperatures were passed unchanged into a one-row logits chunk:

E   RuntimeError: output with shape [1, 11] doesn't match the broadcast shape [7, 11]

41 failed, 23 passed in 2.86s

With the fix, the same test is included in:

64 passed in 2.47s

The TP=2 smoke uses a real NCCL process group and compares each rank's output against an independent per-original-row scalar reference. It covers FP32/BF16, scalar and per-row parameters, CPU/CUDA parameter placement, active subsets, no active rows, partial chunks, and padded vocabulary shards.

The baseline fails on both ranks with:

RuntimeError: output with shape [1, 5] doesn't match the broadcast shape [7, 5]

The fixed version completes all 180 combinations on each rank. The baseline stops at its first failing per-row case; it does not complete the matrix.

To reproduce the complete comparison, extract the attachment and run this from an RL-Kernel Git clone, selecting two free GPUs:

RLK_PY=python CUDA_VISIBLE_DEVICES=0,1 \
  bash /path/to/extracted/reproduce-pr-464.sh

The wrapper fetches the PR, exports the pinned revisions into temporary directories, and runs both pytest and the included tp2_sampling_smoke.py. It does not change the checked-out branch or install dependencies.

I also extracted the original ZIP and executed this wrapper end to end in a separate checkout. It reproduced the same results; that additional session is included under attachment-replay/. The source snapshots were checked against the pinned Git revisions.

These results validate the helper-level fix, including real TP=2 execution; they do not claim full-model training validation or a performance improvement.pr-464-evidence-v2.zip

@Flink-ddd Flink-ddd left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM, Thanks for adding the before and after logs and TP=2 results. This addresses my review.

@maxiaosong1124 maxiaosong1124 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

LGTM!

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

bug Something isn't working

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants