feat(jit): forward sub-part geometry env vars to the JIT - #1
Open
whn09 wants to merge 5 commits into
Open
Conversation
Infrastructure the unordered dispatch/combine kernels build on: a unified GIN context/signal layout with per-peer barrier signal indexing (gin_resource_alloc, qp_mapping), the NCCL device-comm setup reconciled with upstream's runtime-version probe, and the elastic buffer / Python plumbing that resolves the QP budget at construction. One GIN context supplies one QP; the context count (default 11, explicit via num_allocated_qps within [2, 17]) sets the per-context indexed-signal budget that bounds the per-channel part count and the ScaleOut rank count. Also the ptx/layout/comm helpers these paths need. Co-authored-by: Vladimir Aerov <vaerov@amazon.com> Signed-off-by: Xuan Jiang <xuanj@amazon.com>
Split each channel's token stream into parts and sub-parts, each sent as one batched put that carries an in-band header (iteration, token count, continuation flag) and completes a per-part counting signal. The receiver treats the signal purely as a completion count and validates each batch through its header, so correctness holds no matter what order the puts land in. Front-load a smaller first part to cut time-to-first-forward, share one signal across a part's sub-parts to stay inside the per-channel signal budget, and record a per-(token, k) recv-slot map that combine later returns partials through. Co-authored-by: Vladimir Aerov <vaerov@amazon.com> Signed-off-by: Xuan Jiang <xuanj@amazon.com>
Pack the scale-out return partials contiguously per channel and send them as batched puts, each fused with a signal add on a shared per-channel accumulator; the receiver gates on the accumulated count instead of on arrival order. A dedicated proxy warp takes put issuance off the data warps' critical path through a shared-memory hand-off ring. The reduce epilogue locates partials through the per-(token, k) recv map recorded at dispatch rather than assuming in-place token slots. Carries upstream's expanded-send weight handling (kDoExpandedSend) through the slot walk. Co-authored-by: Vladimir Aerov <vaerov@amazon.com> Signed-off-by: Xuan Jiang <xuanj@amazon.com>
The hybrid (scale-out) path now carries two kernel pairs. The unordered pair (default, hybrid_dispatch_unordered.cuh / hybrid_combine_unordered.cuh) synchronizes through in-band headers and counting signals, so strong signal is not required, and weak signal is enough. The ordered pair is the upstream implementation, kept untouched under its original names (hybrid_dispatch.cuh / hybrid_combine.cuh), and runs under the upstream communication configuration (exclusive per-channel contexts, depth-1024 rings, barrier-only signal budget, VA/strong signals); it publishes a tail via a trailing signal and requires the backend to support strong signals and VA signals. Selection happens at JIT-generation time via EP_HYBRID_KERNEL: the generated source names the variant's header and kernel, so the JIT cache keys the two apart automatically. The combine reduce epilogue is shared between the two paths via a kOrderedLayout template parameter. Signed-off-by: Xuan Jiang <xuanj@amazon.com>
`hybrid_dispatch_unordered.cuh` gates the sub-part geometry behind `#ifndef` (`EP_NUM_SUB_PARTS` 2, `EP_MIN_SUB_TOKENS` 1, `EP_SM100_MIN_SUB_TOKENS` 15), but nothing in the tree sets those macros, so the only way to try a different split is to edit the header and reinstall. Forward the three names as JIT `-D` flags, following the `EP_NUM_TOPK_IDX_BITS` block immediately above (and its `EP_JIT_EXTRA_FLAGS` TODO). All three are device-only -- no host translation unit reads them -- so a JIT-only define cannot desync host and device sizing. `flags` is part of `kernel_signature`, so changing the env re-JITs instead of serving a cached cubin. Unset => no behaviour change.
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.
Re-post of Xuan-1998/DeepEP#42 against this repo, rebased onto
main(ec623f3) and re-measured on that tree — the old PR's numbers were taken on the pre-refactor tree, so they are not quoted here.What
hybrid_dispatch_unordered.cuhgates its sub-part geometry behind#ifndef:EP_NUM_SUB_PARTSEP_MIN_SUB_TOKENSEP_SM100_MIN_SUB_TOKENSNothing in the tree ever sets those macros, so the only way to try a different split today is to edit the header and reinstall. This forwards the three names as JIT
-Dflags, following theEP_NUM_TOPK_IDX_BITSblock immediately above it (and that block'sEP_JIT_EXTRA_FLAGSTODO).Unset ⇒ byte-identical behaviour.
Why it is safe
kNumMaxTokensPerRank.)flagsis part ofkernel_signature, so changing the env re-JITs instead of silently serving a cached cubin compiled with a different geometry.Verified end-to-end
EP_JIT_PRINT_COMPILER_COMMAND=1 EP_NUM_SUB_PARTS=1puts-DEP_NUM_SUB_PARTS=1on the nvcc line, and the resulting cubin lands in a distinct JIT cache entry.Measured effect of the knob it exposes
tests/elastic/test_ep.py, 2 ×p5en.48xlarge(8×H200 + 16 EFA each), EP8×2 = 16 ranks,--hidden=7168 --num-topk=8 --num-experts=256 --num-sms=12 --allow-hybrid-mode=1 --prefer-overlap-with-compute=0 --test-first-only. 3 reps, variants interleaved within each rep, each variant on its ownEP_JIT_CACHE_DIR. Mean over all 16 ranks then over reps; ± is stdev across reps. All 48 rounds exited 0.EP_NUM_SUB_PARTS=1on its own:EP_NUM_SUB_PARTS=1So on its own the knob is roughly a wash — a small prefill win, a small decode regression. Its value is that it composes: stacked with #2 it takes decode dispatch from 367.0 µs to 166.1 ± 0.4 µs (−54.7%), where that PR alone reaches 239.8 µs (−34.7%). That is the case this PR is really enabling — being able to find such a combination without a rebuild.
This PR is pure plumbing and changes no default, so it is worth taking independently of whether #2 is accepted.