diff --git a/csrc/jit/compiler.hpp b/csrc/jit/compiler.hpp index ad01b3cac..55365d4b9 100644 --- a/csrc/jit/compiler.hpp +++ b/csrc/jit/compiler.hpp @@ -70,6 +70,17 @@ class Compiler { // TODO: make it more general, e.g. `EP_JIT_EXTRA_FLAGS` if (int num_topk_idx_bits = get_env("EP_NUM_TOPK_IDX_BITS", 0); num_topk_idx_bits != 0) flags += fmt::format(" -DEP_NUM_TOPK_IDX_BITS={}", num_topk_idx_bits); + + // Part / sub-part geometry defaults in `hybrid_dispatch_unordered.cuh`. They are + // device-only (no + // host caller reads them), so forwarding them as JIT defines cannot desync host and device, + // 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", + "EP_MIN_TOKENS_PER_PART"}) + if (int v = get_env(name, 0); v != 0) + flags += fmt::format(" -D{}={}", name, v); } virtual ~Compiler() = default; diff --git a/deep_ep/include/deep_ep/impls/hybrid_dispatch_unordered.cuh b/deep_ep/include/deep_ep/impls/hybrid_dispatch_unordered.cuh index 816f0571d..d98fa7903 100644 --- a/deep_ep/include/deep_ep/impls/hybrid_dispatch_unordered.cuh +++ b/deep_ep/include/deep_ep/impls/hybrid_dispatch_unordered.cuh @@ -55,6 +55,20 @@ static constexpr int kMinSubTokensDefault = (EP_MIN_SUB_TOKENS) > 1 ? (EP_MIN_SU #define EP_SM100_MIN_SUB_TOKENS 15 #endif +// Minimum tokens a scale-out part must carry to be worth its own `flush_part` put, i.e. the +// part-level analogue of `kMinSubTokensDefault` above. `compute_part_allocation()` only ever +// caps the part count from ABOVE when the indexed-signal budget is tight, and that budget is +// loosest exactly when a channel holds the fewest tokens, so small-batch shapes settle on +// `kMaxParts` -- the worst end of the axis. Set to 1 to disable the clamp entirely, i.e. to +// restore the previous behaviour exactly (a value of 1 must SHORT-CIRCUIT rather than divide by +// one: `kNumMaxTokensPerChannel / 1` still clamps whenever a channel holds fewer tokens than the +// budget allows parts, which is a different geometry from the old code and not a control). +#ifndef EP_MIN_TOKENS_PER_PART +#define EP_MIN_TOKENS_PER_PART 15 +#endif + +static constexpr int kMinTokensPerPart = (EP_MIN_TOKENS_PER_PART) > 1 ? (EP_MIN_TOKENS_PER_PART) : 1; + template __device__ __host__ __forceinline__ int num_sub_parts_at(const int& part_tokens) { if constexpr (kNumSubParts <= 1) { @@ -140,9 +154,16 @@ template 0), kNumScaleoutWarps), int kNumMaxTokensPerChannel = math::constexpr_ceil_div(kNumMaxTokensPerRank, kNumChannels), + int kNumBudgetParts = gin_alloc::constexpr_num_parts( + kNumGinSignals, kNumSMs, kNumQPs, (kNumNotifyWarps > 0), kNumScaleoutWarps), + // NOTES: the parentheses around the comparison are load-bearing -- an unparenthesized + // `>` inside a template parameter list closes the list instead of comparing (same + // reason `(kNumNotifyWarps > 0)` above is wrapped) + int kNumGeomParts = kMinTokensPerPart <= 1 ? kNumBudgetParts + : ((kNumMaxTokensPerChannel / kMinTokensPerPart > 1) + ? kNumMaxTokensPerChannel / kMinTokensPerPart : 1), + int kNumParts = kNumBudgetParts < kNumGeomParts ? kNumBudgetParts : kNumGeomParts, int kPartSize = math::constexpr_ceil_div(kNumMaxTokensPerChannel, kNumParts), int kBatchSize = kPartSize, int kNumSubParts = kNumSubPartsDefault < kBatchSize ? kNumSubPartsDefault : kBatchSize,