worker: async scheduling for TP nodes (MSTAR_TP_ASYNC_SCHED) - #224
Conversation
NSagan271
left a comment
There was a problem hiding this comment.
I'm going to do a more detailed pass through the worker either tomorrow or the day after, but putting in some comments now as well---overall looks good but some parts could use some adjusting/hardening.
Also, sorry @vasilevklart --- I'll ask you to rebase this once the resource pool refactor is in.
|
|
||
| There is deliberately no commit / cancel message. Once broadcast, a | ||
| speculative head executes on every rank of the group; whether it must be | ||
| voided (allocation failure on N, per-rid failure, a continuing rid that |
There was a problem hiding this comment.
The invariant that a voided turn "is decided by each rank from state that is identical on every rank by construction" is currently true, but I can see it being fragile in a few cases: (a) if one rank shares a KV cache between the TP node and another node, and (b) if the kv cache across the different ranks carries different sequence lengths (I'm not sure if that would bite for context parallelism,for example).
I do agree that the invariant has to hold (and it even has to hold without the speculative scheduling, i.e., in the current main), so I'd do make an effort to make it clear that this is a contract that could foreseeably break and we need to watch out for, and add defensive checks to, e.g., prevent a KV cache from being shared between a TP node and a node outside of the TP group.
There was a problem hiding this comment.
Agreed, in case of (a), _verify_shared_kv_caches_stay_in_one_tp_group (096245a) is refused at warmup now
There was a problem hiding this comment.
In case of (b), while under TP, KV is replicated per rank and under Ulysses SP it is head sharded so per-request page counts are identical across ranks today, rank 0 admission makes same assumption on serial path.. Context parallelism would give ranks different page counts and break both.. Is CP planned? We can add a per-node verdicts mode for this
https://docs.nvidia.com/nemo/megatron-bridge/latest/parallelisms.html#context-parallelism
There was a problem hiding this comment.
Context parallelism is planned in the future, but I don't think it's been fully scoped out. I think since the non-async path would also be broken with CP the fix is out-of-scope for this PR, but it is perhaps an immediate next step.
NSagan271
left a comment
There was a problem hiding this comment.
I got a chance to go through and understand the worker code; I've amended some of my previous comments that were incorrect and added some new suggestions.
Feel free to push back on anything; there's always the chance that I've missed something when reading the code.
…eneric container Naomi (mstar-project#224, worker.py:288): the two fields _tp_nospec_from/_tp_nospec_order plus _note_tp_nospec are a bounded FIFO set with O(1) lookup and nothing worker-specific. Now mstar.utils.containers.RecentSet(maxlen); the follower keeps one RecentSet[int].
|
Thank you very much @NSagan271 for all the comments and catching the bugs, I believe everything should now be addressed.. I reran qwen3omni TP2 A/B on the earlier version of some of the changes and everything was smooth, 107.88 tok/s -> 131.38 tok/s async, rerunning the pushed tree right now and will post it here.. Lmk if there is still anything and I will rebase to #228.. Claude thoroughly reviewed the current tree and did not find anything, mostly raised follow ups rather any changes right now |
|
Different hardware from last time since the H200 node is busy but re-ran the qwen3omni thinker TP2 A/B on the pushed tree (3×H100 on the Laude node).. serial 117.38 → async 135.81 text tok/s, ITL p50 9 → 8 ms, identical greedy output on all probes |
…rs rebuild it during theirs (MSTAR_TP_ASYNC_SCHED)
3173ab4 to
f196af7
Compare
|
Rebased onto #228 primarily with Claude.. qwen3omni thinker tp2 async +32% b1 / +45% b8 / +57% b16 over serial |
… pageable torch.tensor
|
... with the last commit's pos_ids fix, pushed |
NSagan271
left a comment
There was a problem hiding this comment.
Mainly non-blocking comments, but one thing that should be fixed: the _verify_shared_kv_caches_stay_in_one_tp_group behavior seems to have gotten dropped in the rebase; it should go in KVCacheManager.post_warmup_validate
| @@ -1639,6 +1834,211 @@ def _thread_outputs_to_speculative( | |||
There was a problem hiding this comment.
Pre-existing issue, but these three lines should use r instead of rid. I think it's a small enough issue that it can be fixed here.
| return "failed rids" | ||
| return None | ||
|
|
||
| def _await_tp_follow_step( |
There was a problem hiding this comment.
(Note, not really actionable) Claude pointed out that this could loop forever if there is an admit failure on a follower rank, or if a rid is not ready because of KV cache capacity/offloading. This would not happen if the ranks are symmetric, and is also a pre-existing failure mode (it would just hang on a collective in the non-speculative case instead of looping forever), so it's out-of-scope for this PR.
But as we start thinking about adding context parallelism (and just making the system more robust), it's probably time to think of this failure mode. I'll put up an issue for it.
For this PR, might be worth it to have some (conservative) maximum wall clock time that the loop can run for, and then failing the RIDs if it outruns the maximum. Rank 0 would still wedge on a collective, but the failure would propagate to the client.
There was a problem hiding this comment.
Thanks for filing the issue, I would maybe push back on the cap though.. By the time a follower times out, the leader has already committed those rids into N+1 and is sitting in the collective so failing them on the follower alone makes the ranks disagree about state. The client gets a per request error for a process level problem
The fix is potentially on the instance level (collective timeout / TP-group watchdog), which addresses the issue better; the loop already warns every 2s once N is done so the hang is at least visible in the logs. If you are thinking about the cap, instead of per-rid, we can raise into the main handler so N's rids fail as a batch (same path as a forward we raised)
NSagan271
left a comment
There was a problem hiding this comment.
LGTM (confirmed with @vasilevklart offline that the _verify_shared_kv_caches_stay_in_one_tp_group check already happens in Engine.load_model, added in #228)
TP nodes are scheduled serially: the leader builds step N+1 after N finishes, sends
ScheduleTPNode, followers build after it arrives. So build + plan + prepare sit on the critical path every step on every rank._can_speculaterefuses TP nodes because followers can't start anything on their own.With
MSTAR_TP_ASYNC_SCHED=1(default off):pop_ready_ridsall-or-nothing in wire order), then do exactly what the leader does: reserve slot, pre-plan, wait for N, clear or submitTPNoSpeculationmarker per step, and a follower settles that before it post-processes NWhy no cancel:
test/modular/tp_async_sim.pyis a small exhaustive model checker over rank interleavings. Run-ahead + cancel fails it: one rank can drop the batch after a faster rank already ran it, which desyncs the collectives = NCCL hang. Gated commit passes but costs a round trip per step. This protocol launches a spec on a rank only after that rank finished the parent step and applied its own verdict. Passes all scenarios at 2 and 3 ranks, the failing variant is kept as a negative test.Rebased onto resource pools (#228) as a fresh branch off main. The hook points moved (
NodeOutputis gone, a step's verdict isExecutingBatch.admit_error/failed_requests; the per-step TP barrier is gone), the protocol didn't. The KV-cache-group check from the previous round is dropped:Engine.load_modelenforces it now.MSTAR_ENGINE_STEP_SYNCmust stay 0 with this flag (documented): it holds the GPU thread until N drains, which serialises the overlap.Also in here:
tp_seqon batches, fresh rids returned to their ready queues when a spec is cleared (existing single-worker leak),pop_ready_ridstreats unknown rids as not ready, a step whose forward raised drops its head, a startup check that the flag agrees across the instance, env var docs. Per-ridfailed_requestsare assumed symmetric across ranks, same as the serial TP path already assumes.Tested on this tree:
prepare_inputsbuiltpos_idswith a pageabletorch.tensor(..., device=cuda), which stream-syncs every step so prepare(N+1) waited for step N. Withtorch.fullasync reaches 166 / 711 / 1171 tok/s = +32 / +45 / +57 % over serial; serial unchanged (it builds N+1 after N anyway).test_micro_scheduler+ worker suites); thetest/modularfailure set is identical to main on the same machine.ScheduleTPNodegets 3 new defaulted fields.Merges clean into #243. Main's greedy decode is currently nondeterministic request-to-request on this model, so the old byte-identity probe can't gate this round.