From 7d5696b3a1638a7c980e46ed77eb876faad158a2 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Wed, 16 Sep 2026 02:09:40 -0700 Subject: [PATCH] dsv4.1: communication kernels and wrappers (#39653) Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Co-authored-by: Claude Opus 5 (1M context) Co-authored-by: Cursor Co-authored-by: DarkSharpness <2040703891@qq.com> Co-authored-by: DarkSharpness <76582120+DarkSharpness@users.noreply.github.com> Co-authored-by: Ke Bao Co-authored-by: Khoa Pham Co-authored-by: Xiaoyu Zhang <1182563586@qq.com> Co-authored-by: Xiaoyu Zhang Co-authored-by: Yuhao Yang <47235274+yhyang201@users.noreply.github.com> Co-authored-by: Yuwei An Co-authored-by: Zhichen Zeng Co-authored-by: Ziyi Xu --- .../csrc/distributed/all_reduce_fusion.cuh | 737 ++++++++++++++++++ .../csrc/distributed/custom_all_reduce.cuh | 11 +- .../jit/csrc/distributed/nvlink_comm.cuh | 660 ++++++++++++++++ .../include/sgl_kernel/distributed/ptx.cuh | 8 +- .../ops/communication/all_reduce_fusion.py | 231 ++++++ .../ops/communication/all_reduce_mhc.py | 361 +++++++++ .../kernels/ops/communication/nvlink_comm.py | 156 ++++ .../ops/speculative/dspark/sharded_greedy.py | 109 +++ .../device_communicators/vocab_gather.py | 226 ++++++ 9 files changed, 2494 insertions(+), 5 deletions(-) create mode 100644 python/sglang/kernels/jit/csrc/distributed/all_reduce_fusion.cuh create mode 100644 python/sglang/kernels/jit/csrc/distributed/nvlink_comm.cuh create mode 100644 python/sglang/kernels/ops/communication/all_reduce_fusion.py create mode 100644 python/sglang/kernels/ops/communication/all_reduce_mhc.py create mode 100644 python/sglang/kernels/ops/communication/nvlink_comm.py create mode 100644 python/sglang/kernels/ops/speculative/dspark/sharded_greedy.py create mode 100644 python/sglang/srt/distributed/device_communicators/vocab_gather.py diff --git a/python/sglang/kernels/jit/csrc/distributed/all_reduce_fusion.cuh b/python/sglang/kernels/jit/csrc/distributed/all_reduce_fusion.cuh new file mode 100644 index 000000000..c23127a07 --- /dev/null +++ b/python/sglang/kernels/jit/csrc/distributed/all_reduce_fusion.cuh @@ -0,0 +1,737 @@ +// Fused deferred-MoE finalize -> 1shot lamport push all-reduce [-> RMSNorm] +// over the CustomAllReduceV2 push plane, for decode-sized batches (bf16). The +// hidden width, top_k and cluster geometry are template parameters; the +// shared-expert add and the RMSNorm epilogue are optional. +// +// `idx == -1` marks a dropped slot (EP: the token was routed to an expert that +// is not local) and contributes nothing. Accumulation is fp32 and the bf16 +// rounding points are exactly the unfused path's (moe_runner/flashinfer_trtllm.py +// finalize -> `shared.add_(routed)` -> fp32-accumulating bf16 all-reduce in rank +// order), so the kNorm=false result is bit-identical to it. +// +// The rank-local finalize never materializes in global memory: each thread +// computes one 16B vector of it and pushes it straight into every peer's push +// slot with unicast `st.relaxed.sys` stores, so no multicast mapping is needed. +// +// Push-plane protocol (see include/sgl_kernel/distributed/communicator.cuh): +// * every rank owns 2 phases x kWorldSize slots of `slot_bytes`; a round +// uses phase `counter & 1`, producer r writes slot r of every peer, the +// consumer polls its own kWorldSize slots until no +0.0 marker remains, +// reduces, and restores the +0.0 markers before it exits; +// * +0.0 payload words are remapped to -0.0 (numerically identical) so a +// written word is never 0 and `word == 0` means "not arrived yet"; +// * the generic push kernel owns one phase counter per block; this kernel +// uses one per row cluster (flipped by the cluster's leader block after a +// cluster barrier) plus a trailing "bumper" cluster that flips every +// remaining one, so the whole array keeps one parity and both kernel +// families can share the plane; +// * every rank must call with the same num_tokens / hidden / top_k / epilogue: +// slots are addressed by 16B vector index of the [T, hidden] row view. +#include +#include +#include + +#include +#include +#include +#include +#include +#include + +#include +#include + +#include + +#include +#include +#include +#include + +namespace sglang { + +using device::distributed::PushWorkSpace; +using host::distributed::CommunicatorRef; + +/// One 16B staging vector (8 bf16) viewed as the 4 u32 words the lamport marker +/// protocol tests, matching the generic push kernel's `LamportTrait`. +using Lamport = device::distributed::LamportTrait; +using StageVec = device::AlignedVector; + +SGL_DEVICE void barrier_cluster_arrive_relaxed() { + asm volatile("barrier.cluster.arrive.relaxed.aligned;" ::: "memory"); +} + +SGL_DEVICE void barrier_cluster_wait() { + asm volatile("barrier.cluster.wait.aligned;" ::: "memory"); +} + +template +struct MoeFinalizeAllReduceParams { + bf16_t* out; // [num_tokens, kHiddenDim], output-only + const bf16_t* gemm2; // [P, kHiddenDim], permuted / padded rows + const int32_t* idx; // [num_tokens * kTopK], -1 = dropped slot + const WeightT* weights; // [num_tokens, kTopK], scaling already folded in + const bf16_t* shared; // [num_tokens, kHiddenDim] (kHasShared only) + const bf16_t* norm_weight; // [kHiddenDim] (kNorm only) + float norm_eps; // kNorm only + // Caller's promise that everything read before the PDL wait is complete when + // the predecessor merely *triggers*: no all-reduce on this plane right before + // it, and the routing metadata's producers finished (PDL completion is not + // transitive through early-triggering kernels). False (the default) waits first. + bool prefetch_metadata; + uint32_t rank; + uint32_t num_tokens; + uint32_t num_push_counters; // full counter array size (bumper range end) + PushWorkSpace ws; + bf16_t* mhc_out = nullptr; + const bf16_t* residual = nullptr; + const float* post = nullptr; + const float* comb = nullptr; + const float* pre = nullptr; + bf16_t* normalized = nullptr; + fp8_e4m3_t* quantized = nullptr; + uint8_t* scales = nullptr; +}; + +template +SGL_DEVICE void mhc_quant_vec( + const MoeFinalizeAllReduceParams& params, + const StageVec& value, + uint32_t token, + uint32_t hvec) { + using namespace device; + fp32x2_t v[4]; + float amax = 0.0f; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + v[j] = cast(value[j]); + amax = fmaxf(amax, fmaxf(fabsf(v[j].x), fabsf(v[j].y))); + } + amax = fmaxf(amax, __shfl_xor_sync(0xffffffff, amax, 1, 4)); + amax = fmaxf(amax, __shfl_xor_sync(0xffffffff, amax, 2, 4)); + const float normalized = amax * (1.0f / 448.0f); + const uint32_t bits = __float_as_uint(normalized); + const uint32_t exponent = (bits >> 23) & 255; + const uint32_t mantissa = bits & 0x7fffff; + const bool bump = mantissa != 0 && !(exponent == 0 && mantissa <= 0x400000); + const uint32_t sf = normalized <= 0 ? 0 : min(exponent + uint32_t(bump), 254u); + const float inv_scale = __uint_as_float(sf == 0 ? 0 : (254 - sf) << 23); + AlignedVector q; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + q[j] = cast( + fp32x2_t{fminf(fmaxf(v[j].x * inv_scale, -448.0f), 448.0f), fminf(fmaxf(v[j].y * inv_scale, -448.0f), 448.0f)}); + } + q.store(params.quantized + static_cast(token) * kHiddenDim, hvec); + if (hvec % 4 == 0) { + const uint32_t g = hvec / 4; + const uint32_t off = (g / 4) * 512 + ((token % 32) * 4 + (token / 32) % 4) * 4 + g % 4; + params.scales[off] = sf; + } +} + +/// HC=4 post mixing of an already BF16-rounded all-reduce vector, in +/// mhc_post_split_h's order: round comb[0]*residual[0], FMA post*x, then the +/// remaining three residual streams. +template +SGL_DEVICE StageVec mhc_post_vec( + const MoeFinalizeAllReduceParams& params, const StageVec& red, uint32_t token, uint32_t hvec) { + using namespace device; + StageVec residual[4]; + fp32x2_t collapsed[4] = {}; +#pragma unroll + for (uint32_t c = 0; c < 4; ++c) { + residual[c].load(params.residual + (static_cast(token) * 4 + c) * kHiddenDim, hvec); + } +#pragma unroll + for (uint32_t c = 0; c < 4; ++c) { + const float post = params.post[token * 4 + c]; + float comb[4]; +#pragma unroll + for (uint32_t r = 0; r < 4; ++r) + comb[r] = params.comb[token * 16 + r * 4 + c]; + StageVec out; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + const auto x = cast(red[j]); + const auto r0 = cast(residual[0][j]); + fp32x2_t acc{fmaf(post, x.x, __fmul_rn(comb[0], r0.x)), fmaf(post, x.y, __fmul_rn(comb[0], r0.y))}; +#pragma unroll + for (uint32_t r = 1; r < 4; ++r) { + const auto v = cast(residual[r][j]); + acc.x = fmaf(comb[r], v.x, acc.x); + acc.y = fmaf(comb[r], v.y, acc.y); + } + out[j] = cast(acc); + if constexpr (kCollapse) { + const auto rounded = cast(out[j]); + const float pre = params.pre[token * 4 + c]; + collapsed[j].x = fmaf(rounded.x, pre, collapsed[j].x); + collapsed[j].y = fmaf(rounded.y, pre, collapsed[j].y); + } + } + out.store(params.mhc_out + (static_cast(token) * 4 + c) * kHiddenDim, hvec); + } + StageVec result; + if constexpr (kCollapse) { +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) + result[j] = cast(collapsed[j]); + } + return result; +} + +/// Row geometry: one 16B vector per thread, one cluster per row, so the block +/// size follows from the hidden width and the cluster size (the tuning knob). +template +struct RowClusterTrait { + static constexpr uint32_t kRowVecs = kHiddenDim / 8; // 16B vectors per row + static constexpr uint32_t kBlockSize = kRowVecs / kClusterSize; // threads per block + static constexpr uint32_t kNumWarps = kBlockSize / device::kWarpThreads; + static_assert(kHiddenDim % 8 == 0, "hidden must be a whole number of 16B vectors"); + static_assert(1 <= kClusterSize && kClusterSize <= 8, "portable cluster sizes only"); + static_assert(kRowVecs % kClusterSize == 0, "cluster size must divide the row's vector count"); + static_assert(kBlockSize % device::kWarpThreads == 0, "block must be whole warps"); + static_assert(kBlockSize <= 1024, "block too large: raise the cluster size"); +}; + +// --- stage 1: the deferred finalize of one 16B vector ------------------------ +// The shared-expert vector is loaded first so that load is in flight while the +// routing rows and the kTopK gathers are fetched. +template +SGL_DEVICE StageVec +finalize_vec(const MoeFinalizeAllReduceParams& params, uint32_t token, uint32_t hvec) { + using namespace device; + const auto* idx = params.idx + static_cast(token) * kTopK; + const auto* weights = params.weights + static_cast(token) * kTopK; + int32_t rows[kTopK]; + WeightT w[kTopK]; +#pragma unroll + for (uint32_t k = 0; k < kTopK; ++k) { + rows[k] = idx[k]; + w[k] = weights[k]; + } + + // delay PDL wait until here + PDLWaitPrimary(); + + StageVec shared_in; + if constexpr (kHasShared) { + shared_in.load(params.shared + static_cast(token) * kHiddenDim, hvec); + } + + StageVec in[kTopK]; +#pragma unroll + for (uint32_t k = 0; k < kTopK; ++k) { + if (rows[k] >= 0) in[k].load(params.gemm2 + static_cast(rows[k]) * kHiddenDim, hvec); + } + + fp32x2_t acc[4]; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + acc[j] = fp32x2_t{0.0f, 0.0f}; + } + +#pragma unroll + for (uint32_t k = 0; k < kTopK; ++k) { + if (rows[k] < 0) continue; +#if SGL_ARCH_BLACKWELL_OR_GREATER + if constexpr (std::is_same_v) { +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + acc[j].x = math::fma_f32_bf16(in[k][j].x, w[k], acc[j].x); + acc[j].y = math::fma_f32_bf16(in[k][j].y, w[k], acc[j].y); + } + } else +#endif + { + const auto w_fp32 = cast(w[k]); +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + const auto [x, y] = cast(in[k][j]); + acc[j].x = fmaf(x, w_fp32, acc[j].x); + acc[j].y = fmaf(y, w_fp32, acc[j].y); + } + } + } + // Two deliberate roundings -- the routed combine, then the bf16 + bf16 add -- + // keep the staged vector bit-identical to the unfused rank-local result. + StageVec out; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + if constexpr (kHasShared) { + const auto routed = cast(cast(acc[j])); + const auto sh = cast(shared_in[j]); + out[j] = cast(fp32x2_t{routed.x + sh.x, routed.y + sh.y}); + } else { + out[j] = cast(acc[j]); + } + } + return out; +} + +// --- the kernel -------------------------------------------------------------- +// Grid dim3(num_tokens [+ 1], kClusterSize) with the cluster along y, so +// blockIdx.x is the token row (and its phase counter) and blockIdx.y the rank +// inside the cluster. The extra cluster (blockIdx.x == num_tokens) is the +// bumper: it only flips the leftover counters [num_tokens, num_push_counters). +template < + uint32_t kWorldSize, + uint32_t kHiddenDim, + uint32_t kTopK, + uint32_t kClusterSize, + bool kUsePDL, + bool kHasShared, + bool kNorm, + typename WeightT, + bool kMhc = false, + bool kQuant = false> +__global__ __launch_bounds__(RowClusterTrait::kBlockSize) + __cluster_dims__(1, kClusterSize, 1) void moe_finalize_all_reduce_kernel( + const __grid_constant__ MoeFinalizeAllReduceParams params) { + namespace cg = cooperative_groups; + using namespace device; + using T = RowClusterTrait; + constexpr uint32_t kRowVecs = T::kRowVecs; + constexpr uint32_t kBlockSize = T::kBlockSize; + constexpr uint32_t kNumWarps = T::kNumWarps; + + const auto tx = threadIdx.x; + const auto row_idx = blockIdx.x; + const auto cluster_rank = blockIdx.y; + // this thread's vector within a row: cluster rank picks the block's chunk + const auto hvec = cluster_rank * kBlockSize + tx; + + // Reading the epoch before the PDL wait can see a predecessor all-reduce + // mid-flip on this plane; prefetch_metadata defers the wait to finalize_vec. + if (!params.prefetch_metadata) PDLWaitPrimary(); + + if (row_idx == params.num_tokens) { + PDLWaitPrimary(); + if constexpr (kQuant) { + // The SF buffer is padded to 128 rows. Active rows are written by the + // norm epilogue; the existing bumper zeros only the disjoint padding. + for (uint32_t off = hvec; off < (kHiddenDim / 32) * 128; off += kRowVecs) { + const uint32_t swizzled_row = (off % 512) / 4; + const uint32_t row = swizzled_row / 4 + (swizzled_row % 4) * 32; + if (row >= params.num_tokens) params.scales[off] = 0; + } + } + if (cluster_rank == 0) { + const auto epoch = distributed::PushEpoch{params.ws}; + __syncthreads(); + epoch.unsafe_flip_range(row_idx, params.num_push_counters); + } + return PDLTriggerSecondary(); + } + + // this cluster's epoch: the counter at blockIdx.x, one per row cluster + // (every block of the cluster reads the same one) + const auto epoch = distributed::PushEpoch{params.ws}; + const auto r = params.rank; + // my slot (`src = r`) inside every peer's workspace, and every peer's slot + // inside mine (`dst = r`), for this epoch + void* push_ptrs[kWorldSize]; +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + push_ptrs[i] = epoch.slot_ptr(/*dst=*/i, /*src=*/r); + } + + // stage 1: finalize this row's vector in registers and push it to every peer + const auto vid = row_idx * kRowVecs + hvec; + { + auto vec = finalize_vec(params, row_idx, hvec); + Lamport::clear_pos_zero(vec.data()); +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + ptx::st_relaxed_16B(vec, push_ptrs[i], vid); + } + } + + // ensure epoch is consumed, so flipping it won't lead to error + if constexpr (!kNorm) barrier_cluster_arrive_relaxed(); + + // stage 2: poll own slots, reduce across ranks, [norm], write, reset markers + void* poll_ptrs[kWorldSize]; +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + poll_ptrs[i] = epoch.slot_ptr(/*dst=*/r, /*src=*/i); + } + StageVec vec[kWorldSize]; + do { + bool has_zero = false; +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + ptx::ld_relaxed_16B(vec[i], poll_ptrs[i], vid); + } +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + // the producer remapped +0.0 words, so a written word is never 0: + // word == 0 <=> the slot still holds the empty marker + has_zero |= Lamport::has_pos_zero(vec[i].data()); + } + if (!has_zero) break; + } while (true); + + if constexpr (!kNorm) { + const auto red = reduce_vec(vec); + ptx::st_global_16B(red, params.out, vid); + if constexpr (kMhc) mhc_post_vec(params, red, row_idx, hvec); + // ensure epoch is consumed, so flipping it won't lead to error + barrier_cluster_wait(); + } else { + // push to peer + __shared__ float smem_sq[kClusterSize][kNumWarps]; + auto red = reduce_vec(vec); + if constexpr (kMhc) { + ptx::st_global_16B(red, params.out, vid); + red = mhc_post_vec(params, red, row_idx, hvec); + } + StageVec w; + w.load(params.norm_weight, hvec); + const auto cluster = cg::this_cluster(); + fp32x2_t acc[4]; + float sq = 0.0f; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + acc[j] = cast(red[j]); + sq = fmaf(acc[j].x, acc[j].x, sq); + sq = fmaf(acc[j].y, acc[j].y, sq); + } + sq = warp::reduce_sum(sq); + const auto lane = tx % kWarpThreads; + const auto warp = tx / kWarpThreads; + if (lane < kClusterSize) { + *cluster.map_shared_rank(&smem_sq[cluster_rank][warp], lane) = sq; + } + cluster.sync(); + float total = 0.0f; +#pragma unroll + for (uint32_t c = 0; c < kClusterSize; ++c) { +#pragma unroll + for (uint32_t wp = 0; wp < kNumWarps; ++wp) { + total += smem_sq[c][wp]; + } + } + const auto factor = math::rsqrt(total / static_cast(kHiddenDim) + params.norm_eps); + StageVec out; +#pragma unroll + for (uint32_t j = 0; j < 4; ++j) { + const auto [wa, wb] = cast(w[j]); + out[j] = cast(fp32x2_t{acc[j].x * factor * wa, acc[j].y * factor * wb}); + } + ptx::st_global_16B(out, kMhc ? params.normalized : params.out, vid); + if constexpr (kQuant) mhc_quant_vec(params, out, row_idx, hvec); + } + PDLTriggerSecondary(); + + // re-establish the empty markers for the next same-phase round + StageVec zero_vec; + Lamport::fill_pos_zero(zero_vec.data()); +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + ptx::st_global_16B(zero_vec, poll_ptrs[i], vid); + } + + if (cluster_rank == 0) epoch.flip(); +} + +// --- host -------------------------------------------------------------------- + +template < + uint32_t kWorldSize, + uint32_t kHiddenDim, + uint32_t kTopK, + uint32_t kClusterSize, + bool kUsePDL, + typename WeightT, + bool kMhc = false, + bool kQuant = false> +struct MoeFinalizeAllReduceKernel { + private: + static_assert(std::is_same_v || std::is_same_v); + using TensorView = tvm::ffi::TensorView; + using Params = MoeFinalizeAllReduceParams; + using Trait = RowClusterTrait; + + template + static constexpr auto kernel = moe_finalize_all_reduce_kernel< + kWorldSize, + kHiddenDim, + kTopK, + kClusterSize, + kUsePDL, + kHasShared, + kNorm, + WeightT, + kMhc, + kQuant>; + + public: + /// out = [allreduce over ranks of] finalize(gemm2_out, idx, weights) [+ shared] [-> RMSNorm(norm_weight, eps)]. + /// `out` ([T, kHiddenDim] bf16) is output-only. `shared_output` and `norm_weight` + /// select the epilogue at runtime (four kernel instantiations per module). + static void + run(CommunicatorRef ref, + TensorView out, + TensorView gemm2_out, + TensorView permuted_idx, + TensorView expert_weights, + std::optional shared_output, + std::optional norm_weight, + double eps, + bool prefetch_metadata) { + static_assert(!kMhc); + run_impl( + ref, + out, + gemm2_out, + permuted_idx, + expert_weights, + shared_output, + norm_weight, + eps, + prefetch_metadata, + std::nullopt, + std::nullopt, + std::nullopt, + std::nullopt); + } + + /// Finalize + all-reduce + HC=4 post; original reduced output is retained. + static void run_mhc( + CommunicatorRef ref, + TensorView out, + TensorView gemm2_out, + TensorView permuted_idx, + TensorView expert_weights, + std::optional shared_output, + TensorView mhc_out, + TensorView residual, + TensorView post, + TensorView comb) { + static_assert(kMhc); + run_impl( + ref, + out, + gemm2_out, + permuted_idx, + expert_weights, + shared_output, + std::nullopt, + 0.0, + false, + mhc_out, + residual, + post, + comb); + } + + static void run_mhc_norm( + CommunicatorRef ref, + TensorView out, + TensorView gemm2_out, + TensorView permuted_idx, + TensorView expert_weights, + std::optional shared_output, + TensorView mhc_out, + TensorView residual, + TensorView post, + TensorView comb, + TensorView pre, + TensorView norm_weight, + double eps, + TensorView normalized) { + static_assert(kMhc); + run_impl( + ref, + out, + gemm2_out, + permuted_idx, + expert_weights, + shared_output, + norm_weight, + eps, + false, + mhc_out, + residual, + post, + comb, + pre, + normalized); + } + + static void run_mhc_quant( + CommunicatorRef ref, + TensorView out, + TensorView gemm2_out, + TensorView permuted_idx, + TensorView expert_weights, + std::optional shared_output, + TensorView mhc_out, + TensorView residual, + TensorView post, + TensorView comb, + TensorView pre, + TensorView norm_weight, + double eps, + TensorView normalized, + TensorView quantized, + TensorView scales) { + static_assert(kMhc && kQuant); + run_impl( + ref, + out, + gemm2_out, + permuted_idx, + expert_weights, + shared_output, + norm_weight, + eps, + false, + mhc_out, + residual, + post, + comb, + pre, + normalized, + quantized, + scales); + } + + private: + static void run_impl( + CommunicatorRef ref, + TensorView out, + TensorView gemm2_out, + TensorView permuted_idx, + TensorView expert_weights, + std::optional shared_output, + std::optional norm_weight, + double eps, + bool prefetch_metadata, + std::optional mhc_out, + std::optional residual, + std::optional post, + std::optional comb, + std::optional pre = std::nullopt, + std::optional normalized = std::nullopt, + std::optional quantized = std::nullopt, + std::optional scales = std::nullopt) { + using namespace host; + const auto& comm = *ref.get(); + const auto& push = comm.get_push_obj(); + CHECK_HOST(push.world_size == kWorldSize) + << "communicator holds " << push.world_size << " ranks, kernel built for " << kWorldSize; + + auto T = SymbolicSize{"num_tokens"}; + auto P = SymbolicSize{"num_permuted_rows"}; + auto TK = SymbolicSize{"num_expanded"}; + SymbolicDevice device; + device.set_options(); + TensorMatcher({T, kHiddenDim}) + .with_strides({kHiddenDim, 1}) + .with_dtype() + .with_device(device) + .verify(out); + TensorMatcher({P, kHiddenDim}) + .with_strides({kHiddenDim, 1}) + .with_dtype() + .with_device(device) + .verify(gemm2_out); + TensorMatcher({T, kTopK}) + .with_strides({kTopK, 1}) + .with_dtype() + .template with_device(device) + .verify(expert_weights); + TK.set_value(T.unwrap() * kTopK); + TensorMatcher({TK}).with_strides({1}).with_dtype().with_device(device).verify(permuted_idx); + if (shared_output.has_value()) { + TensorMatcher({T, kHiddenDim}) + .with_strides({kHiddenDim, 1}) + .with_dtype() + .with_device(device) + .verify(shared_output.value()); + } + if (norm_weight.has_value()) { + TensorMatcher({kHiddenDim}) + .with_strides({1}) + .with_dtype() + .with_device(device) + .verify(norm_weight.value()); + } + const auto num_tokens = static_cast(T.unwrap()); + if constexpr (kQuant) { + CHECK_HOST(num_tokens <= 8); + CHECK_HOST(norm_weight.has_value()); + TensorMatcher({T, kHiddenDim}).with_dtype().with_device(device).verify(quantized.value()); + TensorMatcher({(kHiddenDim / 32) * 128}) + .with_dtype() + .with_device(device) + .verify(scales.value()); + } + if constexpr (kMhc) { + static_assert(kHiddenDim == 5120); + if (norm_weight.has_value()) { + TensorMatcher({T, 4}).with_dtype().with_device(device).verify(pre.value()); + TensorMatcher({T, kHiddenDim}).with_dtype().with_device(device).verify(normalized.value()); + } + TensorMatcher({T, 4, kHiddenDim}) + .with_dtype() + .with_device(device) + .verify(mhc_out.value()) + .verify(residual.value()); + TensorMatcher({T, 4}).with_dtype().with_device(device).verify(post.value()); + TensorMatcher({T, 4, 4}).with_dtype().with_device(device).verify(comb.value()); + } + CHECK_HOST(num_tokens > 0) << "num_tokens must be positive"; + CHECK_HOST(reinterpret_cast(gemm2_out.data_ptr()) % 16 == 0) << "gemm2_out must be 16B aligned"; + CHECK_HOST(reinterpret_cast(out.data_ptr()) % 16 == 0) << "out must be 16B aligned"; + + // the whole [T, hidden] row view is staged by vector index, so it must fit + // one push slot; the generic push kernel's callers pick the slot size + const int64_t nbytes = num_tokens * int64_t(kHiddenDim) * sizeof(bf16_t); + CHECK_HOST(nbytes <= push.slot_bytes) << "num_tokens * hidden * 2 = " << nbytes << " bytes exceeds the " + << push.slot_bytes << "-byte push slot (reduce the batch or enlarge " + << "max_push_size)"; + // one cluster (and phase counter) per row; the bumper cluster is launched + // only when counters are left over for it to flip + CHECK_HOST(num_tokens <= push.num_blocks) + << "num_tokens = " << num_tokens << " exceeds the " << push.num_blocks << " push phase counters of the plane"; + const uint32_t num_clusters = num_tokens + (num_tokens < push.num_blocks ? 1 : 0); + + const auto params = Params{ + .out = static_cast(out.data_ptr()), + .gemm2 = static_cast(gemm2_out.data_ptr()), + .idx = static_cast(permuted_idx.data_ptr()), + .weights = static_cast(expert_weights.data_ptr()), + .shared = shared_output.has_value() ? static_cast(shared_output.value().data_ptr()) : nullptr, + .norm_weight = norm_weight.has_value() ? static_cast(norm_weight.value().data_ptr()) : nullptr, + .norm_eps = static_cast(eps), + .prefetch_metadata = prefetch_metadata, + .rank = push.rank, + .num_tokens = num_tokens, + .num_push_counters = push.num_blocks, + .ws = push.get_workspace(nbytes), + .mhc_out = mhc_out.has_value() ? static_cast(mhc_out.value().data_ptr()) : nullptr, + .residual = residual.has_value() ? static_cast(residual.value().data_ptr()) : nullptr, + .post = post.has_value() ? static_cast(post.value().data_ptr()) : nullptr, + .comb = comb.has_value() ? static_cast(comb.value().data_ptr()) : nullptr, + .pre = pre.has_value() ? static_cast(pre.value().data_ptr()) : nullptr, + .normalized = normalized.has_value() ? static_cast(normalized.value().data_ptr()) : nullptr, + .quantized = quantized.has_value() ? static_cast(quantized.value().data_ptr()) : nullptr, + .scales = scales.has_value() ? static_cast(scales.value().data_ptr()) : nullptr, + }; + + const auto has_shared = shared_output.has_value(); + const auto has_norm = norm_weight.has_value(); + const auto kern = has_shared ? (has_norm ? kernel : kernel) + : (has_norm ? kernel : kernel); + // __cluster_dims__(1, kClusterSize, 1) is compiled in, so a plain launch + // already forms the clusters along y + LaunchKernel(dim3(num_clusters, kClusterSize), Trait::kBlockSize, out.device()).enable_pdl(kUsePDL)(kern, params); + } +}; + +} // namespace sglang diff --git a/python/sglang/kernels/jit/csrc/distributed/custom_all_reduce.cuh b/python/sglang/kernels/jit/csrc/distributed/custom_all_reduce.cuh index 627218dfd..b048d2cb7 100644 --- a/python/sglang/kernels/jit/csrc/distributed/custom_all_reduce.cuh +++ b/python/sglang/kernels/jit/csrc/distributed/custom_all_reduce.cuh @@ -142,7 +142,16 @@ ALL_REDUCE_KERNEL void all_reduce_1shot_push_kernel(const __grid_constant__ AllR const auto r = params.rank; const auto num_vecs = params.num_vecs; const auto num_threads = blockDim.x * gridDim.x; - const auto global_tid = blockIdx.x * blockDim.x + threadIdx.x; + // Round-robin warps to blocks rather than giving each block a contiguous run. + // The grid is pinned to the counter array, so once `num_vecs` stops filling + // `gridDim * blockDim` a block-major index leaves the tail CTAs with nothing + // to do; that happens over a whole 2x band of sizes, between the point where + // `choose_block_size` gives up on 512 and the point where 1024 threads fill + // the grid again. + const auto warp_in_block = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto global_warp_id = blockIdx.x + gridDim.x * warp_in_block; + const auto global_tid = global_warp_id * kWarpThreads + lane_id; PDLWaitPrimary(); const auto epoch = distributed::PushEpoch{params.ws}; diff --git a/python/sglang/kernels/jit/csrc/distributed/nvlink_comm.cuh b/python/sglang/kernels/jit/csrc/distributed/nvlink_comm.cuh new file mode 100644 index 000000000..9a2ae037d --- /dev/null +++ b/python/sglang/kernels/jit/csrc/distributed/nvlink_comm.cuh @@ -0,0 +1,660 @@ +#include +#include +#include + +#include +#include +#include +#include + +#include + +#include +#include +#include + +#include +#include +#include + +namespace sglang { + +using device::distributed::PushWorkSpace; +using device::distributed::Semaphore; + +// Runtime uint32 division as a multiply-high and a shift (round-up magic, +// exact below 2^31); cuda::fast_mod_div needs a newer CCCL than CUDA 13 bundles. +struct fast_mod_div_u32_t { + uint32_t divisor; + uint32_t magic; + uint32_t shift; + + __host__ explicit fast_mod_div_u32_t(uint32_t d) : divisor(d), magic(0), shift(0) { + if (d > 1) { + const uint32_t log2_ceil = 32 - std::countl_zero(d - 1); + const uint32_t p = 31 + log2_ceil; + magic = static_cast(((uint64_t{1} << p) + d - 1) / d); + shift = p - 32; + } + } + + __device__ friend uint32_t operator/(uint32_t n, const fast_mod_div_u32_t& fd) { + return fd.divisor == 1 ? n : __umulhi(n, fd.magic) >> fd.shift; + } + + __device__ friend uint32_t operator%(uint32_t n, const fast_mod_div_u32_t& fd) { + return n - (n / fd) * fd.divisor; + } +}; + +template +struct NVLinkCommPushParams { + const void* __restrict__ input; + const void* __restrict__ residual; + void* __restrict__ output; + uint32_t dst_offset; // AR = rank slot stride; AG = packed token prefix + uint32_t rank; + uint32_t num_push_vecs; + uint32_t num_poll_vecs; + uint32_t num_vecs_per_token; + // Ragged split, reduce-scatter only: rank r owns `avg + (r < rem)` tokens of + // the input starting at `r * avg + min(r, rem)`. + uint32_t tokens_avg; + uint32_t tokens_rem; + fast_mod_div_u32_t vecs_per_token_div; + PushWorkSpace ws; +}; + +struct NVLinkCommPullParams { + const void* __restrict__ input; + const void* __restrict__ residual; + void* __restrict__ output; + uint32_t num_vecs; + // multicast buffer + uint8_t* input_mc; + uint8_t* output_mc; + Semaphore* sem_local; + Semaphore* sem_mc; + uint32_t rank; + uint32_t world_size; +}; + +template +inline constexpr uint32_t get_poll_group(uint32_t world_size) { + if (world_size <= 8) return world_size; + return kHasResidual ? 6 : 8; +} + +inline constexpr uint32_t kPushCTASize = 1024; // max value +inline constexpr uint32_t kPullCTASize = 512; // max value + +#define PUSH_KERNEL __global__ __launch_bounds__(kPushCTASize, 1) +#define PULL_KERNEL __global__ __launch_bounds__(kPullCTASize, 1) + +enum Primitive { + RS = 0b01, // Reduce-Scatter + AG = 0b10, // All-Gather + AR = RS | AG, // All-Reduce = RS + AG +}; + +template +SGL_DEVICE vec_t reduce_vec(vec_t x, vec_t y) { + vec_t arr[2] = {x, y}; + return device::reduce_vec(arr); +} + +/** + * \brief Layout: + * 1. `AG`/`RS`: each rank push to its own slot + * [rank0] | [rank 1] | [rank 2] | ... + * 2. `AG`: each rank push to a contiguous region + * [rank0, rank1, rank2, ...] + * + * `RS` use swizzle layout for push kernel \n + * `AG` use normal linear layout for push kernel + */ +template +PUSH_KERNEL void nvlink_push_kernel(const __grid_constant__ NVLinkCommPushParams params) { + using namespace device; + enable_smem_spilling(); + constexpr uint32_t kVecSize = 16 / sizeof(T); // 16 bytes per vector + using vec_t = device::AlignedVector, kVecSize / 2>; + using Lamport = distributed::LamportTrait; + constexpr uint32_t kGroup = get_poll_group(kWorldSize); + + // Round-robin warps to blocks: the poll domain is this rank's shard, so a + // block-major index would park all of it on the first few CTAs and idle the + // rest of the SMs. + const auto warp_in_block = threadIdx.x / kWarpThreads; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto global_warp_id = blockIdx.x + gridDim.x * warp_in_block; + const auto global_tid = global_warp_id * kWarpThreads + lane_id; + const auto num_threads = blockDim.x * gridDim.x; + + PDLWaitPrimary(); + const auto epoch = distributed::PushEpoch{params.ws}; + + void* push_ptrs[kWorldSize]; + /// NOTE: broadcast write is only fast when world size is large + if constexpr (kWorldSize < 8 && (kPrim & Primitive::AG)) { +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + // `dst_offset` is a slot stride for the all-reduce but a packed token + // prefix for the gather, whose consumer reads the plane linearly; this + // must stay the same address arithmetic as the multicast branch below. + push_ptrs[i] = static_cast(epoch.slot_ptr(/*dst=*/i)) + params.dst_offset; + } + } + + const auto dst_ptr_mc = params.ws.mc_workspace + params.dst_offset + epoch.slot_offset(); + const auto vpt = params.num_vecs_per_token; + + for (auto vid = global_tid; vid < params.num_push_vecs; vid += num_threads) { + if constexpr (kPrim & Primitive::AG) { + vec_t vec; + vec.load(params.input, vid); + if constexpr (kHasResidual && kPrim == Primitive::AG) { + vec_t res; + res.load(params.residual, vid); + vec = reduce_vec(vec, res); + } + Lamport::clear_pos_zero(vec.data()); + if constexpr (kWorldSize < 8) { +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + ptx::st_relaxed_16B(vec, push_ptrs[i], vid); + } + } else { + ptx::st_multimem_16B(vec, dst_ptr_mc, vid); + } + } else /* reduce-scatter only */ { + const auto token_id = vid / params.vecs_per_token_div; + const auto offset = vid % params.vecs_per_token_div; + // Both by a compile-time constant, so this is a mask and a shift. + const auto dst_rank = token_id % kWorldSize; + const auto dst_token_id = token_id / kWorldSize; + // Round-robin over peers so every link stays busy instead of one congesting + const auto avg_tokens = params.tokens_avg; + const auto rem_tokens = params.tokens_rem; + const auto rank_prefix = dst_rank * avg_tokens + std::min(dst_rank, rem_tokens); + const auto src_token = rank_prefix + dst_token_id; + vec_t vec; + vec.load(params.input, src_token * vpt + offset); + const auto dst_ptr = epoch.slot_ptr(dst_rank, params.rank); + Lamport::clear_pos_zero(vec.data()); + ptx::st_relaxed_16B(vec, dst_ptr, dst_token_id * vpt + offset); + } + } + + // Poll addresses are linear in the source rank -- one base, `slot_bytes` + // apart -- so a base plus a vector-index bias replaces a per-peer pointer table. + const auto poll_base = epoch.slot_ptr(params.rank); + const auto slot_vecs = params.ws.slot_bytes / sizeof(vec_t); + vec_t pos_zero_vec; + Lamport::fill_pos_zero(pos_zero_vec.data()); + PDLTriggerSecondary(); + + for (auto vid = global_tid; vid < params.num_poll_vecs; vid += num_threads) { + if constexpr (kPrim & Primitive::RS) { + constexpr uint32_t kNumPairs = kVecSize / 2; + vec_t out_vec; + + if constexpr (kGroup >= kWorldSize) { + vec_t vec[kWorldSize + kHasResidual]; + if constexpr (kHasResidual) vec[kWorldSize].load(params.residual, vid); + do { + bool has_zero = false; +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + ptx::ld_relaxed_16B(vec[i], poll_base, i * slot_vecs + vid); + } +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + has_zero |= Lamport::has_pos_zero(vec[i].data()); + } + if (!has_zero) break; + } while (true); + out_vec = reduce_vec(vec); +#pragma unroll + for (uint32_t i = 0; i < kWorldSize; ++i) { + ptx::st_global_16B(pos_zero_vec, poll_base, i * slot_vecs + vid); + } + } else /* > 1 group: divide into chunks */ { + fp32x2_t acc[kNumPairs]; + constexpr uint32_t kNumGroups = div_ceil(kWorldSize, kGroup); + vec_t vec[kGroup]; + vec_t res; + +#pragma unroll + for (uint32_t g = 0; g < kNumGroups; ++g) { + const auto for_each = [&](auto&& fn) { +#pragma unroll + for (uint32_t j = 0; j < kGroup; ++j) { + const auto i = g * kGroup + j; + if (i >= kWorldSize) continue; + fn(i, j); + } + }; + + // Loaded a group early so the fetch overlaps the last poll; it is + // folded into the accumulator once the groups are done. + if constexpr (kHasResidual) { + if (g + 1 == kNumGroups) res.load(params.residual, vid); + } + + do { + bool has_zero = false; + for_each([&](uint32_t i, uint32_t j) { + // load all the vectors + ptx::ld_relaxed_16B(vec[j], poll_base, i * slot_vecs + vid); + }); + for_each([&](uint32_t, uint32_t j) { + // check for zeros + has_zero |= Lamport::has_pos_zero(vec[j].data()); + }); + if (!has_zero) break; + } while (true); + + for_each([&](uint32_t i, uint32_t j) { +#pragma unroll + for (uint32_t k = 0; k < kNumPairs; ++k) { + const auto [x, y] = cast(vec[j][k]); + acc[k].x = i == 0 ? x : acc[k].x + x; + acc[k].y = i == 0 ? y : acc[k].y + y; + } + ptx::st_global_16B(pos_zero_vec, poll_base, i * slot_vecs + vid); + }); + } + if constexpr (kHasResidual) { +#pragma unroll + for (uint32_t k = 0; k < kNumPairs; ++k) { + const auto [x, y] = cast(res[k]); + acc[k].x += x; + acc[k].y += y; + } + } +#pragma unroll + for (uint32_t k = 0; k < kNumPairs; ++k) { + out_vec[k] = cast>(acc[k]); + } + } + + out_vec.store(params.output, vid); + } else /* all-gather only */ { + vec_t vec; + do { + ptx::ld_relaxed_16B(vec, poll_base, vid); + } while (Lamport::has_pos_zero(vec.data())); + vec.store(params.output, vid); + ptx::st_global_16B(pos_zero_vec, poll_base, vid); + } + } + + __syncthreads(); + epoch.flip(); +} + +template +PULL_KERNEL void nvlink_pull_kernel(const __grid_constant__ NVLinkCommPullParams params) { + using namespace device; + constexpr uint32_t kVecSize = 16 / sizeof(T); // 16 bytes per vector + using vec_t = device::AlignedVector, kVecSize / 2>; + constexpr uint32_t kNumWarpVecs = kPullUnroll * kWarpThreads; + + // Round-robin chunks to blocks: the global warp index runs block-fastest, so + // neighbouring chunks are driven by different CTAs. + const auto warp_in_block = threadIdx.x / kWarpThreads; + const auto global_warp_id = blockIdx.x + gridDim.x * warp_in_block; + const auto lane_id = threadIdx.x % kWarpThreads; + const auto num_warps = gridDim.x * (kPullCTASize / kWarpThreads); + + PDLWaitPrimary(); + const auto barrier = distributed::McBarrier{params.sem_local, params.sem_mc, params.world_size, 2}; + barrier.arrive_relaxed(/*n=*/0); + __syncthreads(); + + const auto num_whole_chunks = params.num_vecs / kNumWarpVecs; + // warp uniform unrolled path, 0 predicate + for (auto chunk = global_warp_id; chunk < num_whole_chunks; chunk += num_warps) { + vec_t vecs[kPullUnroll]; + const auto base = chunk * kNumWarpVecs + lane_id; + +#pragma unroll + for (uint32_t i = 0; i < kPullUnroll; ++i) { + const auto vid = base + i * kWarpThreads; + if constexpr (kPrim & Primitive::RS) { + ptx::ld_multimem_16B(vecs[i], params.input_mc, vid); + } else { + ptx::ld_global_16B(vecs[i], params.input, vid); + } + } + + if constexpr (kHasResidual) { + vec_t residuals[kPullUnroll]; +#pragma unroll + for (uint32_t i = 0; i < kPullUnroll; ++i) { + residuals[i].load(params.residual, base + i * kWarpThreads); + } +#pragma unroll + for (uint32_t i = 0; i < kPullUnroll; ++i) { + vecs[i] = reduce_vec(vecs[i], residuals[i]); + } + } + +#pragma unroll + for (uint32_t i = 0; i < kPullUnroll; ++i) { + if constexpr (kPrim & Primitive::AG) { + ptx::st_multimem_16B(vecs[i], params.output_mc, base + i * kWarpThreads); + } else { + ptx::st_global_16B(vecs[i], params.output, base + i * kWarpThreads); + } + } + } + + const auto chunk_offset = num_whole_chunks * kNumWarpVecs; + const auto global_tid = global_warp_id * kWarpThreads + lane_id; + const auto global_threads = num_warps * kWarpThreads; + for (auto vid = chunk_offset + global_tid; vid < params.num_vecs; vid += global_threads) { + vec_t vec; + if constexpr (kPrim & Primitive::RS) { + ptx::ld_multimem_16B(vec, params.input_mc, vid); + } else { + ptx::ld_global_16B(vec, params.input, vid); + } + if constexpr (kHasResidual) { + vec_t res; + res.load(params.residual, vid); + vec = reduce_vec(vec, res); + } + if constexpr (kPrim & Primitive::AG) { + ptx::st_multimem_16B(vec, params.output_mc, vid); + } else { + ptx::st_global_16B(vec, params.output, vid); + } + } + + PDLTriggerSecondary(); + __syncthreads(); + if constexpr (kPrim & Primitive::AG) { + barrier.arrive_rel_acq(/*n=*/1); + } else { // no store multimem, only local store + barrier.arrive_relaxed(/*n=*/1); + } +} + +/// Block size for the push kernel: the smallest that still spreads the work +/// over every SM, capped at the launch bound. +inline auto choose_push_block_size(uint32_t num_vecs) -> uint32_t { + static const uint32_t kNumSM = [] { + int device = 0; + CHECK_CUDA(cudaGetDevice(&device)); + return host::runtime::get_sm_count(device); + }(); + for (const uint32_t block_size : {128u, 256u, 384u, 512u}) { + if (host::div_ceil(num_vecs, block_size) <= kNumSM) return block_size; + } + return 1024u; +} + +template +struct NVLinkComm { + private: + using TensorView = tvm::ffi::TensorView; + using PushPlaneObj = host::distributed::PushPlaneObj; + using PullPlaneObj = host::distributed::PullPlaneObj; + using CommunicatorObj = host::distributed::CommunicatorObj; + using CommunicatorRef = host::distributed::CommunicatorRef; + static constexpr uint32_t kVecBytes = 16; + static constexpr uint32_t kVecSize = kVecBytes / sizeof(T); + + public: + struct RouteInfo { + uint32_t prefix_tokens; // exclusive prefix sum + uint32_t num_rank_tokens; // current rank + }; + + static RouteInfo get_routing(uint32_t num_tokens, uint32_t rank, uint32_t world_size) { + const auto avg = num_tokens / world_size; + const auto rem = num_tokens % world_size; + return {rank * avg + std::min(rank, rem), avg + (rank < rem ? 1 : 0)}; + } + + struct HostParams { + int64_t hidden_size; + DLDevice device; + }; + + /// \brief Base pointer of the residual, shifted onto this rank's slice when + /// the caller hands over the whole tensor. + /// + /// A shard-shaped residual passes straight through; a full tensor is sliced + /// via `get_routing`, not a uniform stride, so ragged splits keep working. + static const void* get_residual_ptr( + const tvm::ffi::Optional& residual, + uint32_t domain_tokens, + uint32_t total_tokens, + uint32_t prefix_bytes) { + if (!residual.has_value()) return nullptr; + const auto tokens = static_cast(residual.value().size(0)); + const auto* base = static_cast(residual.value().data_ptr()); + if (tokens == domain_tokens) return base; + CHECK_HOST(tokens == total_tokens) << "residual has " << tokens << " tokens, expected " << domain_tokens + << " (this rank's shard) or " << total_tokens << " (the whole tensor)"; + return base + prefix_bytes; + } + + static HostParams check_params( + const TensorView in, + const TensorView out, + const tvm::ffi::Optional& residual = {}, + host::DebugInfo info = {}) { + using namespace host; + auto D = SymbolicSize{"hidden_size"}; + auto device_ = SymbolicDevice{}; + auto dtype_ = SymbolicDType{}; + if constexpr (!std::is_same_v) dtype_.set_options(); + device_.set_options(); + TensorMatcher({-1, D}) // + .with_dtype(dtype_) + .with_device(device_) + .verify(in, info); + TensorMatcher({-1, D}) // + .with_dtype(dtype_) + .with_device(device_) + .verify(out, info); + if (residual.has_value()) { + TensorMatcher({-1, D}) // + .with_dtype(dtype_) + .with_device(device_) + .verify(residual.value(), info); + } + return {D.unwrap(), device_.unwrap()}; + } + + private: + template + static void run_push( + const PushPlaneObj& push, + const TensorView in, + const TensorView out, + const tvm::ffi::Optional residual) { + CHECK_HOST(push.world_size == kWorldSize) << push.world_size << " != " << kWorldSize; + const auto [hidden_size, device] = check_params(in, out, residual); + const auto rank = push.rank; + const auto num_vecs_per_token = static_cast(hidden_size / kVecSize); + const auto num_push_vecs = static_cast(in.numel() / kVecSize); + const auto num_poll_vecs = static_cast(out.numel() / kVecSize); + const auto slot_bytes = static_cast(push.slot_bytes); + const auto out_nbytes = static_cast(out.numel() * sizeof(T)); + const auto num_tokens = static_cast(in.size(0)); + const auto out_tokens = static_cast(out.size(0)); + const auto total_tokens = kPrim == Primitive::AG ? out_tokens : num_tokens; + const auto routing = get_routing(total_tokens, rank, kWorldSize); + + uint32_t dst_offset = 0; + if constexpr (kPrim == Primitive::AR) { + CHECK_HOST(num_tokens == out_tokens); + CHECK_HOST(out_nbytes <= push.slot_bytes); + dst_offset = static_cast(rank * slot_bytes); + } else if constexpr (kPrim == Primitive::RS) { + CHECK_HOST(out_tokens == routing.num_rank_tokens); + CHECK_HOST(out_nbytes <= push.slot_bytes); + // dst_offset is not used for this case + } else { + static_assert(kPrim == Primitive::AG); + CHECK_HOST(num_tokens == routing.num_rank_tokens); + CHECK_HOST(out_nbytes <= slot_bytes * kWorldSize); + dst_offset = routing.prefix_tokens * static_cast(num_vecs_per_token * kVecBytes); + } + + // Slice the whole plane: the kernel reaches every slot, not just this rank's. + const auto block_size = choose_push_block_size(std::max(num_push_vecs, num_poll_vecs)); + CHECK_HOST(num_vecs_per_token > 0) << "fast div-mod rejects a zero divisor"; + const auto in_tokens_total = static_cast(in.size(0)); + // The all-reduce reduces the whole tensor on every rank; the other two work + // on this rank's shard, so a full-length residual is sliced. + const auto residual_domain = kPrim == Primitive::AR ? total_tokens : routing.num_rank_tokens; + const auto residual_ptr = get_residual_ptr( + residual, + residual_domain, + total_tokens, + routing.prefix_tokens * static_cast(num_vecs_per_token * kVecBytes)); + const auto params = NVLinkCommPushParams{ + .input = in.data_ptr(), + .residual = residual_ptr, + .output = out.data_ptr(), + .dst_offset = dst_offset, + .rank = rank, + .num_push_vecs = num_push_vecs, + .num_poll_vecs = num_poll_vecs, + .num_vecs_per_token = num_vecs_per_token, + .tokens_avg = in_tokens_total / kWorldSize, + .tokens_rem = in_tokens_total % kWorldSize, + .vecs_per_token_div = fast_mod_div_u32_t{num_vecs_per_token}, + .ws = push.get_workspace(/*size=*/0), + }; + const auto kernel = residual.has_value() ? nvlink_push_kernel + : nvlink_push_kernel; + host::LaunchKernel(push.num_blocks, block_size, device).enable_pdl(kUsePDL)(kernel, params); + } + + template + static void run_pull( + const PullPlaneObj& pull, + const TensorView in, + const TensorView out, + const tvm::ffi::Optional residual, + uintptr_t in_mc_ptr, + uintptr_t out_mc_ptr, + uint32_t num_blocks_hint) { + CHECK_HOST(pull.mc_semaphore != nullptr); + const auto [hidden_size, device] = check_params(in, out, residual); + const auto rank = pull.rank; + const auto world_size = pull.world_size; + const auto num_tokens = static_cast(in.size(0)); + const auto out_tokens = static_cast(out.size(0)); + const auto total_tokens = kPrim == Primitive::AG ? out_tokens : num_tokens; + const auto routing = get_routing(total_tokens, rank, world_size); + const auto num_vecs_per_token = static_cast(hidden_size / kVecSize); + const auto bytes_per_token = static_cast(num_vecs_per_token * kVecBytes); + const auto prefix_bytes = static_cast(routing.prefix_tokens * bytes_per_token); + + // 0 = no hint, autotune; > 0 always use hint but clip to upper bound + if constexpr (kPrim == Primitive::AR) { + CHECK_HOST(num_tokens == out_tokens && in_mc_ptr != 0 && out_mc_ptr != 0); + in_mc_ptr += prefix_bytes; + out_mc_ptr += prefix_bytes; + if (num_blocks_hint == 0) num_blocks_hint = host::div_ceil(256u, kPullUnroll * world_size); + } else if constexpr (kPrim == Primitive::RS) { + CHECK_HOST(out_tokens == routing.num_rank_tokens && in_mc_ptr != 0); + in_mc_ptr += prefix_bytes; + if (num_blocks_hint == 0) num_blocks_hint = pull.num_blocks; // use all the blocks for RS + } else { + static_assert(kPrim == Primitive::AG); + CHECK_HOST(num_tokens == routing.num_rank_tokens && out_mc_ptr != 0); + out_mc_ptr += prefix_bytes; + if (num_blocks_hint == 0) num_blocks_hint = host::div_ceil(128u, kPullUnroll * world_size); + } + /// NOTE: hard limit upper bound is `pull.num_blocks` + num_blocks_hint = std::min(num_blocks_hint, pull.num_blocks); + + // Every pull primitive works on this rank's shard, so a full-length + // residual is sliced onto it. + const auto residual_ptr = get_residual_ptr(residual, routing.num_rank_tokens, total_tokens, prefix_bytes); + const auto params = NVLinkCommPullParams{ + .input = in.data_ptr(), + .residual = residual_ptr, + .output = out.data_ptr(), + .num_vecs = static_cast(routing.num_rank_tokens * num_vecs_per_token), + .input_mc = std::bit_cast(in_mc_ptr), + .output_mc = std::bit_cast(out_mc_ptr), + .sem_local = pull.semaphores[rank], + .sem_mc = pull.mc_semaphore, + .rank = rank, + .world_size = pull.world_size, + }; + + /// NOTE: the final num_blocks resolution must be world unified, otherwise may deadlock + const auto max_vecs_in_world = host::div_ceil(total_tokens, world_size) * num_vecs_per_token; + const auto max_num_blocks = host::div_ceil(max_vecs_in_world, kPullUnroll * kPullCTASize); + const auto num_blocks = std::max(1u, std::min(max_num_blocks, num_blocks_hint)); + const auto kernel = residual.has_value() ? nvlink_pull_kernel + : nvlink_pull_kernel; + host::LaunchKernel(num_blocks, kPullCTASize, device).enable_pdl(kUsePDL)(kernel, params); + } + + public: + // specialized for each world size + template + static void + all_reduce_push(CommunicatorRef comm, TensorView in, TensorView out, tvm::ffi::Optional residual) { + return run_push(comm->get_push_obj(), in, out, residual); + } + template + static void + all_gather_push(CommunicatorRef comm, TensorView in, TensorView out, tvm::ffi::Optional residual) { + return run_push(comm->get_push_obj(), in, out, residual); + } + template + static void + reduce_scatter_push(CommunicatorRef comm, TensorView in, TensorView out, tvm::ffi::Optional residual) { + return run_push(comm->get_push_obj(), in, out, residual); + } + + // only compile once for each world size + template + static void all_reduce_pull( + CommunicatorRef comm, // only pull is needed + TensorView in, + TensorView out, + tvm::ffi::Optional residual, + int64_t in_mc_ptr, + int64_t out_mc_ptr, + uint32_t num_blocks_hint) { + return run_pull( + comm->get_pull_obj(), in, out, residual, in_mc_ptr, out_mc_ptr, num_blocks_hint); + } + template + static void all_gather_pull( + CommunicatorRef comm, // only pull is needed + TensorView in, + TensorView out, + tvm::ffi::Optional residual, + int64_t out_mc_ptr, + uint32_t num_blocks_hint) { + return run_pull( + comm->get_pull_obj(), in, out, residual, 0, out_mc_ptr, num_blocks_hint); + } + template + static void reduce_scatter_pull( + CommunicatorRef comm, // only pull is needed + TensorView in, + TensorView out, + tvm::ffi::Optional residual, + int64_t in_mc_ptr, + uint32_t num_blocks_hint) { + return run_pull(comm->get_pull_obj(), in, out, residual, in_mc_ptr, 0, num_blocks_hint); + } +}; + +} // namespace sglang diff --git a/python/sglang/kernels/jit/include/sgl_kernel/distributed/ptx.cuh b/python/sglang/kernels/jit/include/sgl_kernel/distributed/ptx.cuh index 0b2983da9..ae232abb5 100644 --- a/python/sglang/kernels/jit/include/sgl_kernel/distributed/ptx.cuh +++ b/python/sglang/kernels/jit/include/sgl_kernel/distributed/ptx.cuh @@ -97,7 +97,7 @@ SGL_DEVICE void ld_multimem_16B(V& x, const void* mc_addr, int64_t vec_offset) { mc_addr = static_cast(mc_addr) + vec_offset * 16; if constexpr (std::is_same_v>) { float4 val; - asm volatile("multimem.ld_reduce.weak.add.v4.f32 {%0, %1, %2, %3}, [%4];" + asm volatile("multimem.ld_reduce.weak.global.add.v4.f32 {%0, %1, %2, %3}, [%4];" : "=f"(val.x), "=f"(val.y), "=f"(val.z), "=f"(val.w) : "l"(mc_addr)); x = *reinterpret_cast(&val); @@ -107,12 +107,12 @@ SGL_DEVICE void ld_multimem_16B(V& x, const void* mc_addr, int64_t vec_offset) { // rejects .f32 ("=f") destinations with "Arguments mismatch". uint4 val; if constexpr (std::is_same_v>) { - asm volatile("multimem.ld_reduce.weak.add.acc::f32.v4.f16x2 {%0, %1, %2, %3}, [%4];" + asm volatile("multimem.ld_reduce.weak.global.add.acc::f32.v4.f16x2 {%0, %1, %2, %3}, [%4];" : "=r"(val.x), "=r"(val.y), "=r"(val.z), "=r"(val.w) : "l"(mc_addr)); } else { static_assert(std::is_same_v>); // 4x bf16x2 - asm volatile("multimem.ld_reduce.weak.add.acc::f32.v4.bf16x2 {%0, %1, %2, %3}, [%4];" + asm volatile("multimem.ld_reduce.weak.global.add.acc::f32.v4.bf16x2 {%0, %1, %2, %3}, [%4];" : "=r"(val.x), "=r"(val.y), "=r"(val.z), "=r"(val.w) : "l"(mc_addr)); } @@ -150,7 +150,7 @@ SGL_DEVICE void st_multimem_16B(const V& x, void* mc_addr, int64_t vec_offset) { static_assert(alignof(V) == 16 && sizeof(V) == 16); const auto val = *reinterpret_cast(&x); mc_addr = static_cast(mc_addr) + vec_offset * 16; - asm volatile("multimem.st.weak.v4.f32 [%4], {%0, %1, %2, %3};" + asm volatile("multimem.st.weak.global.v4.f32 [%4], {%0, %1, %2, %3};" : : "f"(val.x), "f"(val.y), "f"(val.z), "f"(val.w), "l"(mc_addr)); #else diff --git a/python/sglang/kernels/ops/communication/all_reduce_fusion.py b/python/sglang/kernels/ops/communication/all_reduce_fusion.py new file mode 100644 index 000000000..47eff6e12 --- /dev/null +++ b/python/sglang/kernels/ops/communication/all_reduce_fusion.py @@ -0,0 +1,231 @@ +"""Fused deferred-MoE finalize + 1shot push all-reduce [+ RMSNorm] (bf16). + +One entry point, :func:`moe_finalize_all_reduce`, over +``csrc/distributed/all_reduce_fusion.cuh``:: + + out[t] = allreduce( sum_k expert_weights[t, k] * gemm2_out[idx[t*top_k + k]] + (+ shared_output[t]) ) # then, optionally, + out[t] = out[t] * rsqrt(mean(out[t]^2) + eps) * norm_weight + +The rank-local finalize (the trtllm-gen ``do_finalize=False`` triple, see +``moe_runner/flashinfer_trtllm.py``) is computed in registers and pushed +straight into every peer's CustomAllReduceV2 push slot, so it never +materializes; ``idx == -1`` slots (EP: non-local expert) contribute nothing. +Small-batch only: the whole ``[T, hidden]`` bf16 row view must fit one push +slot (checked C++-side; :func:`fits_push_slot` lets callers pre-check). + +Needs :func:`register_comm` once per process (the CustomAllReduceV2 +``Communicator``); the ops key on ``world_size`` alone. +""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.kernels.jit.utils import ( + cache_once, + get_jit_cuda_arch, + is_arch_support_pdl, + is_hip_runtime, + load_jit, + make_cpp_args, +) +from sglang.srt.utils.custom_op import register_custom_op + +if TYPE_CHECKING: + from tvm_ffi.module import Module + + from sglang.kernels.ops.communication.all_reduce import Communicator + + +# Storage plane: the CustomAllReduceV2 Communicator (push plane only) + +_COMM_MAP: dict[int, Communicator] = {} + + +def register_comm(comm: Communicator) -> None: + """Register the CustomAllReduceV2 communicator whose push plane the fused + kernel stages through. ``world_size`` is the whole key, so at most one + communicator per size may be registered in a process. + """ + prev = _COMM_MAP.get(comm.world_size) + assert prev is None or prev is comm, ( + f"a different communicator is already registered for world_size=" + f"{comm.world_size}; these ops key only on world_size, so two groups of " + f"the same size cannot coexist in one process" + ) + _COMM_MAP[comm.world_size] = comm + + +def get_registered_comm(world_size: int) -> Optional[Communicator]: + return _COMM_MAP.get(world_size) + + +# Geometry + +_VEC_ELEMS = 8 # bf16 per 16B vector = per thread +_MAX_CLUSTER_SIZE = 8 # portable cluster size limit + + +def valid_cluster_sizes(hidden_dim: int) -> list[int]: + """Cluster sizes the kernel can be built for at this hidden width: whole + 16B vectors per row, whole warps per block, <= 1024 threads, <= 8 blocks.""" + if hidden_dim % _VEC_ELEMS != 0: + return [] + row_vecs = hidden_dim // _VEC_ELEMS + return [ + c + for c in range(1, _MAX_CLUSTER_SIZE + 1) + if row_vecs % c == 0 and (row_vecs // c) % 32 == 0 and row_vecs // c <= 1024 + ] + + +@cache_once +def default_cluster_size(hidden_dim: int) -> int: + if hidden_dim % 1024 == 0 and hidden_dim <= 8192: + return hidden_dim // 1024 + if hidden_dim % 512 == 0 and hidden_dim <= 3584: + return hidden_dim // 512 + candidates = valid_cluster_sizes(hidden_dim) + if not candidates: + raise ValueError( + f"hidden_dim={hidden_dim} has no valid cluster geometry (needs a " + f"multiple of {_VEC_ELEMS * 32} bf16)" + ) + # closest to 128 threads per block, larger block on ties + return min(candidates, key=lambda c: (abs(hidden_dim // _VEC_ELEMS // c - 128), c)) + + +def fits_push_slot(max_push_size: int, num_tokens: int, hidden_dim: int) -> bool: + """Whether a ``[num_tokens, hidden_dim]`` bf16 row view fits one push slot + (``CustomAllReduceV2.max_push_size``).""" + return 0 < num_tokens * hidden_dim * 2 <= max_push_size + + +# JIT module: one per (world_size, hidden_dim, top_k, cluster_size, weight_dtype); the +# shared-add and norm variants are compiled into it and picked at call time. + + +def require_cluster_launch_arch() -> None: + if is_hip_runtime() or get_jit_cuda_arch().major < 9: + raise RuntimeError( + "fused all-reduce cluster kernels require CUDA SM90 or newer" + ) + + +@cache_once +def _jit_module( + world_size: int, + hidden_dim: int, + top_k: int, + cluster_size: int, + weight_dtype: torch.dtype, +) -> Module: + require_cluster_launch_arch() + assert cluster_size in valid_cluster_sizes(hidden_dim), ( + f"cluster_size={cluster_size} is not valid for hidden_dim={hidden_dim}; " + f"choose from {valid_cluster_sizes(hidden_dim)}" + ) + args = make_cpp_args( + world_size, hidden_dim, top_k, cluster_size, is_arch_support_pdl(), weight_dtype + ) + return load_jit( + "moe_finalize_all_reduce", + *args, + cuda_files=["distributed/all_reduce_fusion.cuh"], + cuda_wrappers=[("run", f"MoeFinalizeAllReduceKernel<{args}>::run")], + ) + + +@register_custom_op(mutates_args=["out"]) +def _moe_finalize_all_reduce_op( + world_size: int, + hidden_dim: int, + top_k: int, + cluster_size: int, + out: torch.Tensor, + gemm2_out: torch.Tensor, + expanded_idx_to_permuted_idx: torch.Tensor, + expert_weights: torch.Tensor, + shared_output: Optional[torch.Tensor], + norm_weight: Optional[torch.Tensor], + norm_eps: float, + prefetch_metadata: bool, +) -> None: + comm = _COMM_MAP.get(world_size) + assert comm is not None, ( + f"no communicator registered for world_size={world_size}; call " + "all_reduce_fusion.register_comm(comm.obj) first" + ) + _jit_module(world_size, hidden_dim, top_k, cluster_size, expert_weights.dtype).run( + comm, + out, + gemm2_out, + expanded_idx_to_permuted_idx, + expert_weights, + shared_output, + norm_weight, + norm_eps, + prefetch_metadata, + ) + + +def moe_finalize_all_reduce( + gemm2_out: torch.Tensor, + expanded_idx_to_permuted_idx: torch.Tensor, + expert_weights: torch.Tensor, + top_k: int, + shared_output: Optional[torch.Tensor] = None, + norm_weight: Optional[torch.Tensor] = None, + norm_eps: Optional[float] = None, + *, + world_size: int, + hidden_dim: int, + cluster_size: Optional[int] = None, + prefetch_metadata: bool = False, +) -> torch.Tensor: + """Deferred MoE finalize [+ shared add] -> 1shot push all-reduce [-> RMSNorm]. + + :param gemm2_out: ``[P, hidden_dim]`` bf16, trtllm-gen permuted / padded rows. + :param expanded_idx_to_permuted_idx: ``[T * top_k]`` int32, ``-1`` = dropped slot. + :param expert_weights: ``[T, top_k]`` bf16 or fp32; any routed scaling factor is + already folded in (nothing is rescaled here). + :param shared_output: optional ``[T, hidden_dim]`` bf16 added before the reduce. + :param norm_weight: optional ``[hidden_dim]`` bf16 RMSNorm weight; with + ``norm_eps`` it turns on the fused norm epilogue. + :param prefetch_metadata: read the plane's phase counter and the routing + metadata before the PDL wait; valid only when the + preceding kernel is not an all-reduce on the same + plane and the producers of + ``expanded_idx_to_permuted_idx`` / + ``expert_weights`` are complete. Defaults to False. + :returns: a new ``[T, hidden_dim]`` bf16 tensor (not in place). + """ + num_tokens = expert_weights.shape[0] + assert expert_weights.dtype in (torch.bfloat16, torch.float32) + assert expert_weights.shape[1] == top_k, (expert_weights.shape, top_k) + assert (norm_weight is None) == (norm_eps is None), ( + "norm_weight and norm_eps must be given together" + ) + out = torch.empty( + num_tokens, hidden_dim, dtype=torch.bfloat16, device=gemm2_out.device + ) + if num_tokens == 0: # nothing staged: no phase flip on any rank, stays in step + return out + _moe_finalize_all_reduce_op( + world_size, + hidden_dim, + top_k, + cluster_size or default_cluster_size(hidden_dim), + out, + gemm2_out, + expanded_idx_to_permuted_idx, + expert_weights, + shared_output, + norm_weight, + float(norm_eps) if norm_eps is not None else 0.0, + prefetch_metadata, + ) + return out diff --git a/python/sglang/kernels/ops/communication/all_reduce_mhc.py b/python/sglang/kernels/ops/communication/all_reduce_mhc.py new file mode 100644 index 000000000..ad32c99a4 --- /dev/null +++ b/python/sglang/kernels/ops/communication/all_reduce_mhc.py @@ -0,0 +1,361 @@ +"""MoE finalize and TP all-reduce with an HC=4 post-mixing epilogue.""" + +from typing import Optional, Tuple + +import torch + +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) +from sglang.kernels.ops.communication.all_reduce_fusion import ( + default_cluster_size, + get_registered_comm, + require_cluster_launch_arch, +) +from sglang.srt.utils.custom_op import register_custom_op + +# The kernel is built for this width only (static_assert in all_reduce_fusion.cuh). +_MHC_HIDDEN_DIM = 5120 + + +@cache_once +def _jit_mhc_module(world_size, top_k, cluster_size, weight_dtype): + require_cluster_launch_arch() + args = make_cpp_args( + world_size, + _MHC_HIDDEN_DIM, + top_k, + cluster_size, + is_arch_support_pdl(), + weight_dtype, + True, + ) + return load_jit( + "moe_finalize_all_reduce_mhc", + *args, + cuda_files=["distributed/all_reduce_fusion.cuh"], + cuda_wrappers=[ + ("run", f"MoeFinalizeAllReduceKernel<{args}>::run_mhc"), + ("run_norm", f"MoeFinalizeAllReduceKernel<{args}>::run_mhc_norm"), + ], + ) + + +@register_custom_op(mutates_args=["out", "mhc_out"]) +def _moe_finalize_all_reduce_mhc_op( + world_size: int, + top_k: int, + cluster_size: int, + out: torch.Tensor, + mhc_out: torch.Tensor, + gemm2: torch.Tensor, + idx: torch.Tensor, + weights: torch.Tensor, + shared: Optional[torch.Tensor], + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, +) -> None: + comm = get_registered_comm(world_size) + assert comm is not None + _jit_mhc_module(world_size, top_k, cluster_size, weights.dtype).run( + comm, + out, + gemm2, + idx, + weights, + shared, + mhc_out, + residual, + post, + comb, + ) + + +def moe_finalize_all_reduce_mhc( + gemm2: torch.Tensor, + idx: torch.Tensor, + weights: torch.Tensor, + top_k: int, + shared: Optional[torch.Tensor], + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, + *, + world_size: int, + cluster_size: Optional[int] = None, +) -> Tuple[torch.Tensor, torch.Tensor]: + """Deferred MoE finalize -> push all-reduce -> HC=4 post mixing. + + Same finalize inputs as :func:`all_reduce_fusion.moe_finalize_all_reduce`; + ``residual`` is ``[T, 4, hidden]`` bf16, ``post`` ``[T, 4]`` and ``comb`` + ``[T, 4, 4]`` fp32. Returns ``(reduced [T, hidden], mhc_out [T, 4, hidden])``. + """ + out = torch.empty( + (weights.shape[0], _MHC_HIDDEN_DIM), dtype=torch.bfloat16, device=gemm2.device + ) + mhc_out = torch.empty_like(residual) + if weights.shape[0]: + _moe_finalize_all_reduce_mhc_op( + world_size, + top_k, + cluster_size or default_cluster_size(_MHC_HIDDEN_DIM), + out, + mhc_out, + gemm2, + idx, + weights, + shared, + residual, + post, + comb, + ) + return out, mhc_out + + +@register_custom_op(mutates_args=["out", "mhc_out", "normalized"]) +def _moe_finalize_all_reduce_mhc_norm_op( + world_size: int, + top_k: int, + cluster_size: int, + out: torch.Tensor, + mhc_out: torch.Tensor, + normalized: torch.Tensor, + gemm2: torch.Tensor, + idx: torch.Tensor, + weights: torch.Tensor, + shared: Optional[torch.Tensor], + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, + pre: torch.Tensor, + norm_weight: torch.Tensor, + eps: float, +) -> None: + comm = get_registered_comm(world_size) + assert comm is not None + _jit_mhc_module(world_size, top_k, cluster_size, weights.dtype).run_norm( + comm, + out, + gemm2, + idx, + weights, + shared, + mhc_out, + residual, + post, + comb, + pre, + norm_weight, + eps, + normalized, + ) + + +def moe_finalize_all_reduce_mhc_norm( + gemm2: torch.Tensor, + idx: torch.Tensor, + weights: torch.Tensor, + top_k: int, + shared: Optional[torch.Tensor], + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, + pre: torch.Tensor, + norm_weight: torch.Tensor, + eps: float, + *, + world_size: int, + cluster_size: Optional[int] = None, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """:func:`moe_finalize_all_reduce_mhc` plus the HC=4 pre-collapse + (``pre`` ``[T, 4]`` fp32) and RMSNorm of the collapsed row. Returns + ``(reduced, mhc_out, normalized [T, hidden])``. + """ + out = torch.empty( + (weights.shape[0], _MHC_HIDDEN_DIM), dtype=torch.bfloat16, device=gemm2.device + ) + mhc_out = torch.empty_like(residual) + normalized = torch.empty_like(out) + if weights.shape[0]: + _moe_finalize_all_reduce_mhc_norm_op( + world_size, + top_k, + cluster_size or default_cluster_size(_MHC_HIDDEN_DIM), + out, + mhc_out, + normalized, + gemm2, + idx, + weights, + shared, + residual, + post, + comb, + pre, + norm_weight, + eps, + ) + return out, mhc_out, normalized + + +@cache_once +def _identity_routing(rows, device): + return ( + torch.arange(rows, device=device, dtype=torch.int32), + torch.ones(rows, 1, device=device, dtype=torch.float32), + ) + + +def all_reduce_mhc_norm( + x: torch.Tensor, + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, + pre: torch.Tensor, + norm_weight: torch.Tensor, + eps: float, + *, + world_size: int, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Plain all-reduce of ``x`` with the mHC + norm epilogue: the finalize + kernel driven with identity routing (top_k = 1, unit weights).""" + idx, weights = _identity_routing(x.shape[0], x.device) + return moe_finalize_all_reduce_mhc_norm( + x, + idx, + weights, + 1, + None, + residual, + post, + comb, + pre, + norm_weight, + eps, + world_size=world_size, + ) + + +@cache_once +def _jit_mhc_quant_module(world_size, top_k, cluster_size, weight_dtype): + require_cluster_launch_arch() + args = make_cpp_args( + world_size, + _MHC_HIDDEN_DIM, + top_k, + cluster_size, + is_arch_support_pdl(), + weight_dtype, + True, + True, + ) + return load_jit( + "moe_finalize_all_reduce_mhc_quant", + *args, + cuda_files=["distributed/all_reduce_fusion.cuh"], + cuda_wrappers=[("run", f"MoeFinalizeAllReduceKernel<{args}>::run_mhc_quant")], + ) + + +@register_custom_op( + mutates_args=["out", "mhc_out", "normalized", "quantized", "scales"] +) +def _moe_finalize_all_reduce_mhc_quant_op( + world_size: int, + top_k: int, + cluster_size: int, + out: torch.Tensor, + mhc_out: torch.Tensor, + normalized: torch.Tensor, + quantized: torch.Tensor, + scales: torch.Tensor, + gemm2: torch.Tensor, + idx: torch.Tensor, + weights: torch.Tensor, + shared: Optional[torch.Tensor], + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, + pre: torch.Tensor, + norm_weight: torch.Tensor, + eps: float, +) -> None: + comm = get_registered_comm(world_size) + assert comm is not None + _jit_mhc_quant_module(world_size, top_k, cluster_size, weights.dtype).run( + comm, + out, + gemm2, + idx, + weights, + shared, + mhc_out, + residual, + post, + comb, + pre, + norm_weight, + eps, + normalized, + quantized, + scales, + ) + + +def moe_finalize_all_reduce_mhc_quant( + gemm2: torch.Tensor, + idx: torch.Tensor, + weights: torch.Tensor, + top_k: int, + shared: Optional[torch.Tensor], + residual: torch.Tensor, + post: torch.Tensor, + comb: torch.Tensor, + pre: torch.Tensor, + norm_weight: torch.Tensor, + eps: float, + *, + world_size: int, + cluster_size: Optional[int] = None, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """:func:`moe_finalize_all_reduce_mhc_norm` plus fp8 e4m3 quantization of + the normalized row with ue8m0 group scales (rows <= 8). Returns + ``(reduced, mhc_out, normalized, quantized, scales)``. + """ + rows = weights.shape[0] + assert 0 < rows <= 8 + out = torch.empty( + (rows, _MHC_HIDDEN_DIM), dtype=torch.bfloat16, device=gemm2.device + ) + mhc_out = torch.empty_like(residual) + normalized = torch.empty_like(out) + quantized = torch.empty_like(out, dtype=torch.float8_e4m3fn) + # ue8m0 scale layout: one byte per 32-wide group, rows padded to 128 + scales = torch.empty( + (_MHC_HIDDEN_DIM // 32) * 128, device=gemm2.device, dtype=torch.uint8 + ) + _moe_finalize_all_reduce_mhc_quant_op( + world_size, + top_k, + cluster_size or default_cluster_size(_MHC_HIDDEN_DIM), + out, + mhc_out, + normalized, + quantized, + scales, + gemm2, + idx, + weights, + shared, + residual, + post, + comb, + pre, + norm_weight, + eps, + ) + return out, mhc_out, normalized, quantized, scales diff --git a/python/sglang/kernels/ops/communication/nvlink_comm.py b/python/sglang/kernels/ops/communication/nvlink_comm.py new file mode 100644 index 000000000..6a7a0f485 --- /dev/null +++ b/python/sglang/kernels/ops/communication/nvlink_comm.py @@ -0,0 +1,156 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, Final + +import torch + +from sglang.kernels.jit.utils import ( + cache_once, + is_arch_support_pdl, + load_jit, + make_cpp_args, +) + +if TYPE_CHECKING: + from tvm_ffi import Module + + from sglang.kernels.ops.communication.all_reduce import Communicator + + +_PRIMITIVES: Final = ["all_reduce", "all_gather", "reduce_scatter"] + + +def get_multicast_ptr(tensor: torch.Tensor) -> int: + """Multicast alias of a symmetric-memory tensor. Collective on first call; + torch caches the handle per allocation, so repeats stay cheap. + """ + from torch._C._distributed_c10d import _SymmetricMemory + + ptr = _SymmetricMemory.rendezvous(tensor).multicast_ptr + assert ptr != 0, "tensor has no multicast alias; was it allocated p2p?" + return ptr + + +@cache_once +def _jit_pull_module(dtype: torch.dtype, num_unroll: int) -> Module: + args = make_cpp_args(dtype, is_arch_support_pdl()) + return load_jit( + "nvlink_comm_pull", + *args, + f"unroll{num_unroll}", + cuda_files=["distributed/nvlink_comm.cuh"], + cuda_wrappers=[ + (n, f"NVLinkComm<{args}>::{n}_pull<{num_unroll}>") for n in _PRIMITIVES + ], + ) + + +@cache_once +def _jit_push_module(dtype: torch.dtype, world_size: int) -> Module: + args = make_cpp_args(dtype, is_arch_support_pdl()) + return load_jit( + "nvlink_comm_push", + *args, + f"world{world_size}", + cuda_files=["distributed/nvlink_comm.cuh"], + cuda_wrappers=[ + (n, f"NVLinkComm<{args}>::{n}_push<{world_size}>") for n in _PRIMITIVES + ], + ) + + +# `residual` on any of these is folded into the reduction; it may be shaped like +# this rank's shard or like the whole tensor, of which this rank's slice is taken. +def all_reduce_push( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + residual: torch.Tensor | None = None, +) -> None: + _jit_push_module(input.dtype, comm.world_size).all_reduce( + comm, input, output, residual + ) + + +def all_gather_push( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + residual: torch.Tensor | None = None, +) -> None: + _jit_push_module(input.dtype, comm.world_size).all_gather( + comm, input, output, residual + ) + + +def reduce_scatter_push( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + residual: torch.Tensor | None = None, +) -> None: + _jit_push_module(input.dtype, comm.world_size).reduce_scatter( + comm, input, output, residual + ) + + +def all_reduce_pull( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + residual: torch.Tensor | None = None, + *, + in_mc_ptr: int = 0, + out_mc_ptr: int = 0, + num_unroll=4, + num_blocks_hint: int = 0, +) -> None: + _jit_pull_module(input.dtype, num_unroll).all_reduce( + comm, + input, + output, + residual, + in_mc_ptr or get_multicast_ptr(input), + out_mc_ptr or get_multicast_ptr(output), + num_blocks_hint, + ) + + +def all_gather_pull( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + residual: torch.Tensor | None = None, + *, + out_mc_ptr: int = 0, + num_unroll=4, + num_blocks_hint: int = 0, +) -> None: + _jit_pull_module(input.dtype, num_unroll).all_gather( + comm, + input, + output, + residual, + out_mc_ptr or get_multicast_ptr(output), + num_blocks_hint, + ) + + +def reduce_scatter_pull( + comm: Communicator, + input: torch.Tensor, + output: torch.Tensor, + residual: torch.Tensor | None = None, + *, + in_mc_ptr: int = 0, + num_unroll=4, + num_blocks_hint: int = 0, +) -> None: + _jit_pull_module(input.dtype, num_unroll).reduce_scatter( + comm, + input, + output, + residual, + in_mc_ptr or get_multicast_ptr(input), + num_blocks_hint, + ) diff --git a/python/sglang/kernels/ops/speculative/dspark/sharded_greedy.py b/python/sglang/kernels/ops/speculative/dspark/sharded_greedy.py new file mode 100644 index 000000000..b86bd4468 --- /dev/null +++ b/python/sglang/kernels/ops/speculative/dspark/sharded_greedy.py @@ -0,0 +1,109 @@ +"""Exact greedy selection from TP-local base logits and rounded Markov bias.""" + +import torch +import triton +import triton.language as tl + + +@triton.jit +def _sharded_greedy_partial_kernel( + B, + X, + P, + BS: tl.constexpr, + XS: tl.constexpr, + WIDTH: tl.constexpr, + OFFSET: tl.constexpr, + PARTS: tl.constexpr, + BLOCK: tl.constexpr, +): + row, part = tl.program_id(0), tl.program_id(1) + i = part * BLOCK + tl.arange(0, BLOCK) + valid = i < WIDTH + bias = tl.load(B + row * BS + i, valid, 0).to(tl.float32) + base = tl.load(X + row * XS + i, valid, 0).to(tl.float32) + value = base + bias + nan = valid & (value != value) + has_nan = tl.sum(nan.to(tl.int32), 0) > 0 + maximum = tl.max(tl.where(valid & ~nan, value, -float("inf")), 0) + wins = tl.where(has_nan, nan, valid & (value == maximum)) + idx = tl.min(tl.where(wins, i + OFFSET, 2147483647), 0) + maximum = tl.where(has_nan, float("nan"), maximum) + tl.store(P + (row * PARTS + part) * 2, maximum) + tl.store(P + (row * PARTS + part) * 2 + 1, idx.to(tl.float32, bitcast=True)) + + +@triton.jit +def _sharded_greedy_finish_kernel( + P, + OUT, + BS: tl.constexpr, + PARTS: tl.constexpr, + WORLD: tl.constexpr, + BLOCK: tl.constexpr, +): + row = tl.program_id(0) + i = tl.arange(0, BLOCK) + valid = i < PARTS * WORLD + rank, part = i // PARTS, i % PARTS + offset = ((rank * BS + row) * PARTS + part) * 2 + value = tl.load(P + offset, valid, -float("inf")) + idx = tl.load(P + offset + 1, valid, 0).to(tl.int32, bitcast=True) + # The NVLink push transport changes +0 to -0 as its Lamport sentinel. + # Indices are nonnegative int32 bits, so strip that sign bit to recover ID 0. + idx &= 2147483647 + valid &= idx != 2147483647 + nan = valid & (value != value) + has_nan = tl.sum(nan.to(tl.int32), 0) > 0 + maximum = tl.max(tl.where(valid & ~nan, value, -float("inf")), 0) + wins = tl.where(has_nan, nan, valid & (value == maximum)) + result = tl.min(tl.where(wins, idx, 2147483647), 0) + tl.store(OUT + row, result.to(tl.int64)) + + +def sharded_greedy_step(bias, base_local, *, group, vocab_start, gather=None): + """Fused BuildStepLocal + vocab gather + argmax, without materializing logits. + + Equivalent to argmax of rank-ordered ``build_step_local``/all_gather over the + sharded vocab, excluding padding. The transport carries one (value, + global-index-bits) pair per 4096-wide block of the shard per row; indices + move as bits and are never converted numerically to float. + """ + assert bias.ndim == base_local.ndim == 2 + assert bias.shape[0] == base_local.shape[0] + assert bias.shape[1] <= base_local.shape[1] + assert bias.stride(1) == base_local.stride(1) == 1 + rows, width = bias.shape + block = 4096 + parts = triton.cdiv(base_local.shape[1], block) + assert parts > 0 + partial = torch.empty((rows, parts, 2), device=bias.device, dtype=torch.float32) + _sharded_greedy_partial_kernel[(rows, parts)]( + bias, + base_local, + partial, + bias.stride(0), + base_local.stride(0), + width, + vocab_start, + parts, + block, + num_warps=4, + ) + # The padded partition width fixes the transport shape on all ranks; + # WIDTH masks real entries, including a completely empty final shard. + if gather is not None: + gathered = gather(partial.view(rows, parts * 2)) + else: + gathered = group.all_gather(partial, dim=0) if group.world_size > 1 else partial + result = torch.empty(rows, device=bias.device, dtype=torch.int64) + _sharded_greedy_finish_kernel[(rows,)]( + gathered, + result, + rows, + parts, + group.world_size, + triton.next_power_of_2(parts * group.world_size), + num_warps=4, + ) + return result diff --git a/python/sglang/srt/distributed/device_communicators/vocab_gather.py b/python/sglang/srt/distributed/device_communicators/vocab_gather.py new file mode 100644 index 000000000..4ea0dbcf0 --- /dev/null +++ b/python/sglang/srt/distributed/device_communicators/vocab_gather.py @@ -0,0 +1,226 @@ +"""All-gather of a vocab-parallel row block across a TP group. + +Every implementation takes this rank's ``[rows, local_width]`` slice and returns +``[rows, world_size * local_width]`` with the ranks' slices side by side, the +layout ``GroupCoordinator.all_gather(dim=-1)`` produces. +""" + +from __future__ import annotations + +import logging +from abc import ABC, abstractmethod +from typing import Optional, Tuple + +import torch + +logger = logging.getLogger(__name__) + + +class VocabGather(ABC): + """``[rows, local] -> [rows, world_size * local]``, ranks side by side.""" + + @abstractmethod + def __call__(self, local: torch.Tensor) -> torch.Tensor: ... + + @abstractmethod + def gather_stacked(self, local: torch.Tensor) -> torch.Tensor: + """Gather compact row blocks as [world_size * rows, local_width].""" + ... + + +class LocalVocabGather(VocabGather): + """A group of one: the slice is the whole row.""" + + def __call__(self, local: torch.Tensor) -> torch.Tensor: + return local + + def gather_stacked(self, local: torch.Tensor) -> torch.Tensor: + return local + + +class NcclVocabGather(VocabGather): + """The group coordinator's all_gather along the last dim (NCCL ring).""" + + def __init__(self, group) -> None: + self.group = group + + def __call__(self, local: torch.Tensor) -> torch.Tensor: + return self.group.all_gather(local, dim=-1) + + def gather_stacked(self, local: torch.Tensor) -> torch.Tensor: + return self.group.all_gather(local, dim=0) + + +# Collective: every rank of the group must call this, in the same order, outside +# CUDA-graph capture. The returned multicast alias is 0 when the group has none. +def _alloc_symm( + group, shape: Tuple[int, int], dtype: torch.dtype +) -> Tuple[torch.Tensor, int]: + from torch._C._distributed_c10d import _SymmetricMemory + + # a GroupCoordinator names the allocation by its cpu_group, as + # CustomAllReduceV2 does; a torch process group names it itself + pg = getattr(group, "cpu_group", group) + buf = _SymmetricMemory.empty_strided_p2p( + (shape[0] * shape[1],), + [1], + dtype, + torch.device("cuda", torch.cuda.current_device()), + pg.group_name, + ) + mc_ptr = int(_SymmetricMemory.rendezvous(buf).multicast_ptr) + return buf.view(shape), mc_ptr + + +class NVLinkVocabGather(VocabGather): + """The NVLink collectives on CustomAllReduceV2's multicast plane. + + A slice that fits one slot of the push plane takes the push kernel into a + fresh tensor; a larger one that fits ``pull_out`` takes the pull kernel into + that symmetric-memory output, which is reused every call; anything else goes + to ``fallback``, the NCCL ring. Both kernels gather along the row axis, so + the ranks come back stacked and are transposed into place. ``pull_out`` is + allocated here: the allocation is collective and captured graphs keep its + address. + """ + + def __init__( + self, + *, + ca_comm, + group, + local_width: int, + dtype: torch.dtype, + symm_rows: int, + fallback: VocabGather, + ) -> None: + self.comm = ca_comm.obj + self.world_size = int(group.world_size) + self.slot_bytes = int(ca_comm.max_push_size) + self.fallback = fallback + self.pull_out: Optional[torch.Tensor] = None + self.pull_mc_ptr = 0 + if symm_rows > 0 and self.comm.pull is not None: + out, mc_ptr = _alloc_symm( + group, (self.world_size * symm_rows, local_width), dtype + ) + if mc_ptr != 0: + self.pull_out, self.pull_mc_ptr = out, mc_ptr + logger.info( + "NVLink vocab gather: pull output %s (%d MB)", + tuple(out.shape), + out.numel() * out.element_size() >> 20, + ) + else: + logger.warning("NVLink vocab gather: no multicast alias, pull path off") + + def __call__(self, local: torch.Tensor) -> torch.Tensor: + rows = local.shape[0] + if local.nbytes <= self.slot_bytes: + return self._push(local) + total_rows = self.world_size * rows + if self.pull_out is not None and total_rows <= self.pull_out.shape[0]: + return self._pull(local, self.pull_out[:total_rows]) + return self.fallback(local) + + def gather_stacked(self, local: torch.Tensor) -> torch.Tensor: + # Compact argmax partials need rank-major output, no symmetric pull buffer. + if ( + local.is_contiguous() + and local.shape[1] * local.element_size() % 16 == 0 + and local.nbytes <= self.slot_bytes + ): + return self._push_stacked(local) + return self.fallback.gather_stacked(local) + + def _push(self, local: torch.Tensor) -> torch.Tensor: + return self._unstack(self._push_stacked(local)) + + def _push_stacked(self, local: torch.Tensor) -> torch.Tensor: + from sglang.kernels.ops.communication import nvlink_comm + + rows, width = local.shape + gathered = torch.empty( + (self.world_size * rows, width), dtype=local.dtype, device=local.device + ) + nvlink_comm.all_gather_push(self.comm, local, gathered) + return gathered + + def _pull(self, local: torch.Tensor, out: torch.Tensor) -> torch.Tensor: + from sglang.kernels.ops.communication import nvlink_comm + + nvlink_comm.all_gather_pull(self.comm, local, out, out_mc_ptr=self.pull_mc_ptr) + full = self._unstack(out) + # the transpose copies except at one row, where it would alias the + # shared buffer that the next call overwrites + return full.clone() if full.data_ptr() == out.data_ptr() else full + + def _unstack(self, gathered: torch.Tensor) -> torch.Tensor: + rows = gathered.shape[0] // self.world_size + width = gathered.shape[1] + if rows == 1: + return gathered.view(1, self.world_size * width) + return ( + gathered.view(self.world_size, rows, width) + .transpose(0, 1) + .reshape(rows, self.world_size * width) + ) + + +def _nvlink_ca_comm(group, *, local_width: int, dtype: torch.dtype): + ca_comm = getattr(group, "ca_comm", None) + if ca_comm is None or getattr(ca_comm, "disabled", True): + return None + comm = getattr(ca_comm, "obj", None) + if comm is None or not getattr(ca_comm, "has_multicast", False): + return None + if comm.push is None or comm.world_size != group.world_size: + return None + # the kernels move 16-byte vectors along the row + if local_width % (128 // torch.finfo(dtype).bits) != 0: + return None + return ca_comm + + +def _default_symm_rows() -> int: + try: + from sglang.srt.runtime_context import get_exec, get_schedule + + return int( + get_schedule().max_running_requests + or get_exec().graph.cuda_graph_config.decode.max_bs + or 0 + ) + except Exception: + return 0 + + +def make_vocab_gather( + group, + *, + local_width: int, + dtype: torch.dtype = torch.float32, + prefer_nvlink: bool = True, + symm_rows: Optional[int] = None, +) -> VocabGather: + """The gather for ``group``: local for a group of one, NVLink when the + group's custom all-reduce has a multicast plane (and ``prefer_nvlink``), + the NCCL ring otherwise. ``symm_rows`` is the row capacity of the NVLink + gather's symmetric-memory output for slices past the push slot; None sizes + it for the server's largest batch, 0 leaves those slices to NCCL.""" + if group is None or group.world_size == 1: + return LocalVocabGather() + nccl = NcclVocabGather(group) + if not prefer_nvlink: + return nccl + ca_comm = _nvlink_ca_comm(group, local_width=local_width, dtype=dtype) + if ca_comm is None: + return nccl + return NVLinkVocabGather( + ca_comm=ca_comm, + group=group, + local_width=local_width, + dtype=dtype, + symm_rows=_default_symm_rows() if symm_rows is None else symm_rows, + fallback=nccl, + )