feat(jit): forward sub-part geometry env vars to the JIT - #1
Conversation
efe26c2 to
5118c2e
Compare
|
Rebased onto The old branch showed a 3,000-line diff over 20 files, which was an artefact, not this change. It was cut from Measurement caveat: the numbers in the description were taken on |
5118c2e to
1a08a36
Compare
| // and `flags` is part of `kernel_signature` below, so a change re-JITs rather than | ||
| // reusing a stale cubin. Tuning them per network/arch currently requires editing the | ||
| // header and reinstalling. | ||
| for (const auto& name: {"EP_NUM_SUB_PARTS", "EP_MIN_SUB_TOKENS", "EP_SM100_MIN_SUB_TOKENS"}) |
There was a problem hiding this comment.
Can you also update the readme's env var list: https://github.com/amazon-contributing/DeepEP#environment-variables
There was a problem hiding this comment.
Done in the amended head — added EP_NUM_SUB_PARTS, EP_MIN_SUB_TOKENS, EP_SM100_MIN_SUB_TOKENS to the General list next to EP_NUM_TOPK_IDX_BITS, each noting what 0/unset means and the header default it falls back to. Force-pushed to ca76bfc.
`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. (cherry picked from commit 5118c2e)
1a08a36 to
ca76bfc
Compare
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.
Second architecture: B300 /
sm_103over EFASetup. 2 ×
p6-b300(8×B300 SXM6 + 16 EFA gen-3 each, EFA installer 1.50.0,efa.ko3.3.0g, driver 595.91.07), EP8×2 = 16 ranks, NCCL 2.31.2 GIN withNCCL_GIN_TYPE=5(EFA-GDA), same shape flags as above. 3 reps per cell.On this architecture
EP_NUM_SUB_PARTS=1is neutral, and it does not composethe way it does on H200. Measured on top of #2 (the two are in one image, so the
b300 run isolates the knob stacked, not standalone):
EP_NUM_SUB_PARTS=1So the H200 result — that
EP_NUM_SUB_PARTS=1takes stacked decode dispatch from239.8 µs to 166.1 µs — does not carry over to
sm_103, where #2's clamp isalready the whole effect. That is a point in favour of this PR rather than against
it: the right sub-part split is evidently architecture- and shape-dependent, and
today finding it requires editing a header and reinstalling.
A second, stronger safety argument found while measuring on B300
The claim above that
flagsis part ofkernel_signatureis not just a nicety —it is the only thing that makes a geometry change safe to cache. The other half
of that key,
code, is the generated wrapper; the impl headers are only#included, so their content is not hashed, andsignatureis just the compilerversion. We hit this directly: two images differing only by a header edit, sharing
one
EP_JIT_CACHE_DIR, silently reuse each other's cubins and the edit measures asa no-op.
Changing this geometry through the env, as this PR allows, is therefore strictly
safer than the header-edit-and-reinstall path it replaces: the env lands in
flags, which is hashed, so a variant can never be served another variant'scubin. (Worth a separate issue for the header-hash gap, which is orthogonal to
this PR.)