Skip to content

Support context parallelism in the Tinker/multi-LoRA loss path (#3284) - #8

Open
micahtyong wants to merge 54 commits into
mainfrom
devin/1789704380-tinker-cp-support
Open

micahtyong wants to merge 54 commits into
mainfrom
devin/1789704380-tinker-cp-support

Conversation

@micahtyong

@micahtyong micahtyong commented Sep 18, 2026 •

Copy link
Copy Markdown

Summary

Enable --context-parallel-size > 1 for the Tinker / multi-LoRA loss path (stage 2 of radixark#3284). The Tinker losses now consume the local CP shard of every per-token client vector and report full-response per_datum outputs on every CP rank; validate_multi_lora_args drops its context_parallel_size == 1 assert.

Bug

get_log_probs_and_entropy returns each CP rank's zigzag shard of the response log-probs (and rollout_log_probs is already sliced to match), but tinker_losses.py zipped that shard against the client's full-length advantages / loss_weights / loss_masks, so every Tinker loss failed under CP>1 with a shape mismatch. multi_lora.py turned this into a launch-time rejection.

Changes

loss_hub/tinker_losses.py — slice client vectors to the local shard, the same way the native losses do:

_local_shards(args, batch, values)         # slice_log_prob_with_cp per datum
_response_masks(args, batch, log_probs)    # get_local_response_loss_masks (was: full masks)
cross_entropy:          loss_weights -> _local_shards(...)
is / ppo / cispo / dro: advantages   -> _local_shards(...)
rollout_log_probs, target_tokens: unchanged (already local)

The loss stays a local-shard sum; the existing intra_dp_cp.size scaling in loss_function plus Megatron's DP×CP grad averaging give a gradient identical to CP=1. This was chosen over gathering full-length log-probs (the Lilo workaround) because it matches losses.py and adds no collective on the gradient path.

per_datum must be full-response on every rank (MilesBackend._execute_batch keeps the first worker result per sample_index). New _gather_per_datum_outputs:

logprobs: all_gather_with_cp(log_prob.detach(), ...)             # zigzag reassembled
loss:     stack(detached per-datum losses); all_reduce over cp group if cp.size > 1

Both run on detached tensors, so no gradient crosses the CP group.

A rank whose shards of every datum in the microbatch are empty (e.g. a lone short response) has no supervised tokens; _sum_loss_and_outputs now returns logits[..., :0].sum() there so backward still runs against the graph instead of raising on a constant.

utils/multi_lora.py: remove the CP=1 assert; the PP=1 assert stays.

serve_tinker.py: each CP rank holds an even-sized zigzag pair of chunks per datum and pads its microbatch to tp * data_pad_size_multiplier, so the per-datum token cap is (max_tokens_per_gpu floored to that multiple, made even) * cp_size; CP=1 unchanged.

Validation

tests/fast/.../test_tinker_loss_cp2_consistency.py (new): gloo CP group through the real _target_logprobs + loss_function, asserting vs CP=1: reassembled log-probs, all-reduced loss, per-shard logits gradient, identical per_datum on every rank, and that the batch's client vectors stay full-length (a recomputed forward must not see them sliced). Cases: CP2 × all five TINKER_LOSS_FUNCTIONS × {ragged: response spanning both chunks, odd length, rank-0 shard with no response tokens; empty_rank: one-token response leaving rank 0 with no supervised tokens at all}, plus CP4 × {direct, --recompute-loss-function}. Fails on main, passes here. test_arguments.py gains a multi-LoRA CP 2/4 acceptance case; test_tinker_loss_outputs.py only needed a real ParallelState and qkv_format.

pytest tests/fast/backends/training_utils/loss/test_tinker_loss_cp2_consistency.py \
       tests/fast/backends/training_utils/loss/test_tinker_loss_outputs.py \
       tests/fast/backends/training_utils/loss/test_tinker_explicit_targets.py \
       tests/fast/backends/training_utils/loss/test_loss_cp2_consistency.py \
       tests/fast/utils/test_arguments.py::TestMultiLoRAValidation            # passed
pre-commit run --files <changed files>                                          # passed

E2E: serve_tinker.py on this branch, Qwen3-30B-A3B-Instruct-2507 LoRA (rank 32) GRPO through the Tinker SDK, 64k context, trainer 4×H200 TP2×CP2 (EP4), rollout 4×H200 sglang TP1; 4 steps + save_state → load_state on a fresh client → 1 step. reward_mean 0.46 → 0.46 → 0.60 → 0.65 → 0.70, update_successful:mean = 1.0 every step; datums reach 59k tokens (> --max-tokens-per-gpu 32768), exercising the serve_tinker.py cap change. W&B: raw-miles-cp-native-qwen3_30b_a3b-64k-tp2cp2.

Scope

  • A response with prompt_length == 0 is not handled by the existing get_logits_and_tokens_offset_with_cp / all_gather_with_cp helpers; unreachable on the Tinker path (tokens = model_input + target_tokens[-1:]), so not addressed here.
  • --allgather-cp with multi-LoRA remains rejected in get_batch (unchanged).

Refs radixark#3284.

Link to Devin session: https://modal.devinenterprise.com/sessions/beea8ece75d949e18dacbc2c7f66aa4b
Open in Devin Desktop: https://modal.devinenterprise.com/desktop/session/beea8ece75d949e18dacbc2c7f66aa4b?variant=devin
Requested by: @micahtyong

@devin-ai-integration

Copy link
Copy Markdown

🤖 Devin AI Engineer

I'll be helping with this pull request! Here's what you should know:

✅ I will automatically:

  • Address comments on this PR that start with 'DevinAI' or '@devin'.
  • Look at CI failures and help fix them

Note: I can only respond to comments from users who have write access to this repository.

⚙️ Control Options:

  • Disable automatic comment, CI, and merge conflict monitoring

@devin-ai-integration
devin-ai-integration Bot force-pushed the devin/1789704380-tinker-cp-support branch from 26fa230 to 1dcf79b Compare September 18, 2026 22:32
@devin-ai-integration devin-ai-integration Bot changed the title Support context parallelism in the Tinker/multi-LoRA loss path (radixark/miles#3284) Support context parallelism in the Tinker/multi-LoRA loss path (#3284) Sep 18, 2026
@micahtyong

Copy link
Copy Markdown
Author

/devin review

@devin-ai-integration

Copy link
Copy Markdown

Starting Devin Review.

Devin Review

@devin-ai-integration devin-ai-integration Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

✅ Devin Review: No Issues Found

Devin Review analyzed this PR and found no bugs or issues to report.

Devin Review

@devin-ai-integration
devin-ai-integration Bot changed the base branch from main to devin/tmp-base-refresh September 18, 2026 23:05
@devin-ai-integration
devin-ai-integration Bot changed the base branch from devin/tmp-base-refresh to main September 18, 2026 23:05
@micahtyong

Copy link
Copy Markdown
Author

Upstream PR is here: radixark#3320

We prefer to keep our miles fork up-to-date with the actual miles repo, so I won't merge this in and instead will prefer to merge in the upstream, and then rebase the modal-projects/miles fork afterwards.

nanjiangwill and others added 17 commits September 21, 2026 12:38
…SA indexer (radixark#3296)

Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
…ang checkout (radixark#3588)

Co-authored-by: Cursor Agent <cursoragent@cursor.com>
…s it (radixark#3587)

--fully-async selects FullyAsyncRolloutFn, but naming that class through
--rollout-function-path selected the same producer while leaving the mode off, so
the run skipped every --fully-async check (colocate, partial rollout, legacy
rollout v1, pause mode, multi-LoRA) and train.py's async-driver guard.

Normalize the path spelling into the flag before validation runs. Only the exact
class is recognized; a subclass still passes --fully-async explicitly.
…ion creation (radixark#3585)

Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
guapisolo and others added 24 commits September 22, 2026 20:47
Co-authored-by: Jiajun Li <guapisolo@gmail.com>
Signed-off-by: Yusheng Su <yushengsu.thu@gmail.com>
Co-authored-by: Ethan (Yusheng) Su <yushengsu.thu@gmail.com>
Co-authored-by: Zhiyao Jiang <jessicajiang324@gmail.com>
Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
Co-authored-by: yueming-yuan <yym022502@gmail.com>
…ixark#3661)

Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
…ark#3284)

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
The trainer packs max_tokens_per_gpu * context_parallel_size tokens per
microbatch, so the gateway's per-datum cap must scale the same way or CP>1
rejects sequences the trainer can hold.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
A rank whose zigzag shards of every datum in the microbatch are empty
produced a loss with no grad_fn, so backward raised. Gate on local
response tokens rather than on the datum list and use an empty logits
view for the zero loss; add an empty-rank case to the CP2 test.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
A CP rank holds an even number of tokens per datum (a zigzag chunk pair),
so the per-rank budget must be even before scaling by CP size; the old
formula could admit a datum whose padded per-rank share exceeded
max_tokens_per_gpu. Add a multi-LoRA CP validation test, a CP4 leg with
loss recomputation, and assert the batch's client vectors stay full-length
across a recomputed forward.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
@devin-ai-integration
devin-ai-integration Bot force-pushed the devin/1789704380-tinker-cp-support branch from 522af2a to d2bdc01 Compare September 24, 2026 17:29
micahtyong and others added 5 commits September 25, 2026 13:42
…sses

Dropping the CP==1 assert let --allgather-cp through validation even though
get_batch still rejects it with multi-LoRA, so the gateway died on the first
batch instead of at launch. Also clone the unbind() views so each per-datum
loss pickles only its own storage on the Ray return path.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
- _local_shards -> _local_response_values, _response_masks -> _local_response_masks
  (both now return this CP rank's slice)
- _gather_per_datum_outputs: loss_sums / full_losses / local_loss / full_loss /
  sample_index so the rank-local vs CP-summed values and the datum index are
  unambiguous
- test_tinker_loss_cp2_consistency.py -> test_tinker_loss_cp_consistency.py,
  since it also holds the cp4 test

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
- one-line reason for the empty-view loss on a rank with no response tokens
- note that rollout_log_probs arrives pre-sliced by get_rollout_data
- serve_tinker: keep the padding comment, give the even-share rule its own line

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
all_gather_with_cp issued an autograd all_reduce per datum, so a microbatch
of N datums paid N+1 collectives. Scatter each rank's local logprob shard
into a zero response (positions from slicing an arange with the same CP
helper), concatenate with the per-datum losses, and reduce once. Also drop
the redundant SimpleNamespace parallel-state monkeypatch in the output tests;
make_parallel_state already sets the state loss.py reads.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.