Support context parallelism in the Tinker/multi-LoRA loss path (#3284) - #8
Open
micahtyong wants to merge 54 commits into
Open
micahtyong wants to merge 54 commits into
micahtyong wants to merge 54 commits into
Conversation
🤖 Devin AI EngineerI'll be helping with this pull request! Here's what you should know: ✅ I will automatically:
Note: I can only respond to comments from users who have write access to this repository. ⚙️ Control Options:
|
devin-ai-integration
Bot
force-pushed
the
devin/1789704380-tinker-cp-support
branch
from
September 18, 2026 22:32
26fa230 to
1dcf79b
Compare
Author
|
/devin review |
devin-ai-integration
Bot
changed the base branch from
main
to
devin/tmp-base-refresh
September 18, 2026 23:05
devin-ai-integration
Bot
changed the base branch from
devin/tmp-base-refresh
to
main
September 18, 2026 23:05
…ixark#3319) Co-authored-by: yueming-yuan <yym022502@gmail.com>
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. |
…SA indexer (radixark#3296) Co-authored-by: Claude Fable 5.1 <noreply@anthropic.com>
…ad, trainer-owned adapter (radixark#3194)
…radixark#3291) Signed-off-by: Yusheng Su <yushengsu.thu@gmail.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>
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>
…dixark#3653) Co-authored-by: Claude Opus 5.5 <noreply@anthropic.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
Bot
force-pushed
the
devin/1789704380-tinker-cp-support
branch
from
September 24, 2026 17:29
522af2a to
d2bdc01
Compare
…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>
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Enable
--context-parallel-size > 1for 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-responseper_datumoutputs on every CP rank;validate_multi_lora_argsdrops itscontext_parallel_size == 1assert.Bug
get_log_probs_and_entropyreturns each CP rank's zigzag shard of the response log-probs (androllout_log_probsis already sliced to match), buttinker_losses.pyzipped that shard against the client's full-lengthadvantages/loss_weights/loss_masks, so every Tinker loss failed under CP>1 with a shape mismatch.multi_lora.pyturned 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:The loss stays a local-shard sum; the existing
intra_dp_cp.sizescaling inloss_functionplus 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 matcheslosses.pyand adds no collective on the gradient path.per_datummust be full-response on every rank (MilesBackend._execute_batchkeeps the first worker result persample_index). New_gather_per_datum_outputs: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_outputsnow returnslogits[..., :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 totp * 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, identicalper_datumon every rank, and that the batch's client vectors stay full-length (a recomputed forward must not see them sliced). Cases: CP2 × all fiveTINKER_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 onmain, passes here.test_arguments.pygains a multi-LoRA CP 2/4 acceptance case;test_tinker_loss_outputs.pyonly needed a realParallelStateandqkv_format.E2E:
serve_tinker.pyon 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_stateon a fresh client → 1 step. reward_mean 0.46 → 0.46 → 0.60 → 0.65 → 0.70,update_successful:mean = 1.0every step; datums reach 59k tokens (>--max-tokens-per-gpu 32768), exercising theserve_tinker.pycap change. W&B: raw-miles-cp-native-qwen3_30b_a3b-64k-tp2cp2.Scope
prompt_length == 0is not handled by the existingget_logits_and_tokens_offset_with_cp/all_gather_with_cphelpers; unreachable on the Tinker path (tokens = model_input + target_tokens[-1:]), so not addressed here.--allgather-cpwith multi-LoRA remains rejected inget_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