diff --git a/python/pyproject.toml b/python/pyproject.toml index 508e3855e..1be0cdfd5 100755 --- a/python/pyproject.toml +++ b/python/pyproject.toml @@ -40,6 +40,7 @@ dependencies = [ "numpy", "nvidia-cutlass-dsl[cu13]==4.5.2", "nvidia-ml-py", + "nvidia-mathdx==25.6.0", "openai-harmony==0.0.4", "openai==2.6.1", "orjson", diff --git a/python/sglang/jit_kernel/csrc/lplb/dispatch_probability.cuh b/python/sglang/jit_kernel/csrc/lplb/dispatch_probability.cuh new file mode 100644 index 000000000..749464fbe --- /dev/null +++ b/python/sglang/jit_kernel/csrc/lplb/dispatch_probability.cuh @@ -0,0 +1,122 @@ +// LP probability dispatch kernel: collapse the ~7 torch ops in +// _topk_ids_logical_to_physical_probability into one launch. +// +// Python equivalent: +// +// topk_probs = log2phy_prob[topk_ids] # gather +// row_sums = topk_probs.sum(dim=-1) # reduce +// fallback = (log2phy_map[topk_ids] >= 0).float() # gather + cast +// topk_probs = where(row_sums > 0, topk_probs, fallback) # cmp + select +// chosen = multinomial(topk_probs, 1).flatten() # sample +// out = log2phy_map[topk_ids, chosen] # gather +// +// Each thread handles one (token, slot) in the flattened topk_ids. +// `random_vals` is pre-generated by the caller via torch.rand (one kernel +// launch, ~5 µs); we sample the multinomial via prefix-sum + comparison +// against `random_vals[i] * row_sum`. +// +// Templated on (MAX_COPIES, BLOCK_DIM). MAX_COPIES is small (typically 2-3) +// so the per-row prefix sum unrolls into a few instructions per thread. + +#include +#include + +#include + +#include +#include + +#include + +namespace { + +template +__global__ void dispatch_probability_kernel( + int32_t* __restrict__ out_topk_ids, // (N,) + const int32_t* __restrict__ in_topk_ids, // (N,) + const float* __restrict__ log2phy_prob, // (NUM_LOGICAL, MAX_COPIES) + const int32_t* __restrict__ log2phy_map, // (NUM_LOGICAL, MAX_COPIES) + const float* __restrict__ random_vals, // (N,) + int N) { + const int idx = blockIdx.x * BLOCK_DIM + threadIdx.x; + if (idx >= N) return; + + const int32_t logical_id = in_topk_ids[idx]; + const int32_t* row_map = log2phy_map + logical_id * MAX_COPIES; + const float* row_prob = log2phy_prob + logical_id * MAX_COPIES; + + float probs[MAX_COPIES]; + int32_t maps[MAX_COPIES]; + float row_sum = 0.f; + +#pragma unroll + for (int c = 0; c < MAX_COPIES; c++) { + maps[c] = row_map[c]; + probs[c] = row_prob[c]; + row_sum += probs[c]; + } + + // Fallback: if all LP probs for this row are 0, sample uniformly from the + // valid physical copies (map != -1). + if (row_sum <= 0.f) { + row_sum = 0.f; +#pragma unroll + for (int c = 0; c < MAX_COPIES; c++) { + probs[c] = (maps[c] >= 0) ? 1.0f : 0.0f; + row_sum += probs[c]; + } + } + + // Multinomial sample: smallest c such that cumsum[0..c] > u * row_sum. + // Implemented branch-free for unroll friendliness: chosen accumulates the + // largest index where cumsum[..c] is still <= u, then we add 1 (clamped to + // MAX_COPIES-1 for the all-cumsum-<=-u edge case from float rounding). + const float u = random_vals[idx] * row_sum; + float cum = 0.f; + int chosen = 0; +#pragma unroll + for (int c = 0; c < MAX_COPIES; c++) { + cum += probs[c]; + if (u >= cum) chosen = c + 1; + } + if (chosen >= MAX_COPIES) chosen = MAX_COPIES - 1; + + out_topk_ids[idx] = maps[chosen]; +} + +template +void dispatch_probability( + tvm::ffi::TensorView out_topk_ids, // (N,) int32 + tvm::ffi::TensorView in_topk_ids, // (N,) int32 + tvm::ffi::TensorView log2phy_prob, // (NUM_LOGICAL, MAX_COPIES) float32 + tvm::ffi::TensorView log2phy_map, // (NUM_LOGICAL, MAX_COPIES) int32 + tvm::ffi::TensorView random_vals) { // (N,) float32 + using namespace host; + + SymbolicSize N{"num_topk_entries"}; + SymbolicSize NUM_LOGICAL{"num_logical"}; + SymbolicDevice device_; + + TensorMatcher({N}).with_dtype().with_device(device_).verify(out_topk_ids).verify(in_topk_ids); + TensorMatcher({NUM_LOGICAL, MAX_COPIES}).with_dtype().with_device(device_).verify(log2phy_prob); + TensorMatcher({NUM_LOGICAL, MAX_COPIES}).with_dtype().with_device(device_).verify(log2phy_map); + TensorMatcher({N}).with_dtype().with_device(device_).verify(random_vals); + + const int n = static_cast(N.unwrap()); + const int grid = (n + BLOCK_DIM - 1) / BLOCK_DIM; + const DLDevice device = device_.unwrap(); + + using KernelT = void (*)(int32_t*, const int32_t*, const float*, const int32_t*, const float*, int); + KernelT kernel = dispatch_probability_kernel; + + LaunchKernel(grid, BLOCK_DIM, device)( + kernel, + static_cast(out_topk_ids.data_ptr()), + static_cast(in_topk_ids.data_ptr()), + static_cast(log2phy_prob.data_ptr()), + static_cast(log2phy_map.data_ptr()), + static_cast(random_vals.data_ptr()), + n); +} + +} // namespace diff --git a/python/sglang/jit_kernel/csrc/lplb/ipm.cuh b/python/sglang/jit_kernel/csrc/lplb/ipm.cuh new file mode 100644 index 000000000..c420411c1 --- /dev/null +++ b/python/sglang/jit_kernel/csrc/lplb/ipm.cuh @@ -0,0 +1,275 @@ +// Single-SM Interior Point Method (IPM) LP solver. +// +// Solves min c^T x subject to A x = b, x >= 0 with a barrier method: +// for step in 0..NUM_ITERS: +// ax2 = A * x^2 +// ax2a = ax2 @ A^T (cuBLASDx GEMM) +// ax2c = ax2 @ c (cuBLASDx GEMM) +// d = solve(ax2a, ax2c) (cuSolverDx Cholesky/POSV) +// r = d^T @ A +// d = x * (c - r) +// x *= 1 - 0.999 * d / max(d) +// +// Convergence is checked at the end and the kernel writes 0.5 to every +// element on non-convergence (matches the historical Numba behavior). +// +// Adapted from DeepSeek-AI/LPLB's `minilp.cu`. Templated on (NC, NV, +// BLOCK_DIM, SM_VER, NUM_ITERS) so each unique shape is compiled once via +// sglang's tvm-ffi load_jit cache. + +#include +#include + +#include + +#include +#include + +#include +#include + +namespace { + +template +struct ipm_smem { + float b[NC]; + float a[NC][NV]; + float c[NV]; + float ax2[NC][NV]; + float ax2a[NC][NC]; + float x[NV]; + float ax2c[NC]; + float r[NV]; + float d[NV]; + float alpha; + bool avail_flag; +}; + +// In-place Cholesky factorization a = L L^T (lower triangle), no external +// linkage. Replaces cuSolverDx::posv to keep the kernel self-contained +// under sglang's tvm-ffi load_jit (which uses plain c++ for the final +// link step and so cannot satisfy cuSolverDx's device-link requirement). +// +// For the typical LPLB shape (N <= 32) the algorithm is dwarfed by the +// cuBLASDx GEMMs. +template +__device__ __forceinline__ void cholesky_factor(float a[N][N]) { + const int tid = threadIdx.x; + for (int k = 0; k < N; k++) { + if (tid == 0) { + // Clamp the pivot away from zero before sqrtf. Numerical drift in + // the IPM iterations can push a[k][k] slightly negative on + // otherwise-PSD matrices, which would produce NaN and propagate + // through the rest of the solve. The convergence check at the end + // of the kernel writes 0.5 on non-convergence regardless. + a[k][k] = sqrtf(fmaxf(a[k][k], 1e-12f)); + } + __syncthreads(); + const float pivot = a[k][k]; + for (int i = k + 1 + tid; i < N; i += BLOCK_DIM) { + a[i][k] /= pivot; + } + __syncthreads(); + // Schur complement: a[i][j] -= a[i][k] * a[j][k] for j>k, i>=j + for (int idx = tid; idx < N * N; idx += BLOCK_DIM) { + const int i = idx / N, j = idx % N; + if (j > k && i >= j && i < N) { + a[i][j] -= a[i][k] * a[j][k]; + } + } + __syncthreads(); + } +} + +// Solve L L^T x = b in-place on b, where L is the lower triangle of a +// (filled by `cholesky_factor`). Forward then back substitution; both +// run on a single thread because N is small and the inner loops have +// loop-carried dependencies that don't parallelize cheaply. +template +__device__ __forceinline__ void cholesky_apply(const float a[N][N], float b[N]) { + const int tid = threadIdx.x; + if (tid == 0) { + for (int i = 0; i < N; i++) { + float s = b[i]; + for (int j = 0; j < i; j++) { + s -= a[i][j] * b[j]; + } + b[i] = s / a[i][i]; + } + for (int i = N - 1; i >= 0; i--) { + float s = b[i]; + for (int j = i + 1; j < N; j++) { + s -= a[j][i] * b[j]; + } + b[i] = s / a[i][i]; + } + } + __syncthreads(); +} + +template +__device__ __forceinline__ void cholesky_solve(float a[N][N], float b[N]) { + cholesky_factor(a); + cholesky_apply(a, b); +} + +template +__device__ __forceinline__ void matmul_NT(float* a, float* b, float* c) { + decltype(cublasdx::Size() + cublasdx::Function() + cublasdx::Arrangement() + cublasdx::SM() + cublasdx::Block() + cublasdx::BlockDim())() + .execute(1.f, a, b, 0.f, c); +} + +template +__device__ __forceinline__ void matmul_NN(float* a, float* b, float* c) { + decltype(cublasdx::Size() + cublasdx::Function() + cublasdx::Arrangement() + cublasdx::SM() + cublasdx::Block() + cublasdx::BlockDim())() + .execute(1.f, a, b, 0.f, c); +} + +template +__global__ void ipm_solve_kernel( + float* __restrict__ result, + const float* __restrict__ input_a, + const float* __restrict__ input_b, + const float* __restrict__ input_c) { + using SMem = ipm_smem; + extern __shared__ unsigned char raw_smem[]; + SMem* smem = reinterpret_cast(raw_smem); + + const int tid = threadIdx.x; + const int dim = blockDim.x; + + auto& a = smem->a; + auto& b = smem->b; + auto& c = smem->c; + + // Load A, b, c into shared memory (single block, no grid index). + for (int i = tid; i < NC * NV; i += dim) { + int ic = i / NV, iv = i % NV; + a[ic][iv] = input_a[i]; + } + __syncthreads(); + for (int i = tid; i < NC; i += dim) { + b[i] = input_b[i]; + } + __syncthreads(); + for (int i = tid; i < NV; i += dim) { + c[i] = input_c[i]; + } + __syncthreads(); + + auto& ax2 = smem->ax2; + auto& ax2a = smem->ax2a; + auto& x = smem->x; + auto& ax2c = smem->ax2c; + auto& r = smem->r; + auto& d = smem->d; + auto& alpha = smem->alpha; + // d_max and max_residual are warp-scope reductions; only valid for tid<32. + float d_max = 0.f; + float max_residual = 0.f; + + for (int j = tid; j < NV; j += dim) { + x[j] = 1.f; + } + __syncthreads(); + + for (int step = 0; step < NUM_ITERS; step++) { + for (int ij = tid; ij < NC * NV; ij += dim) { + int i = ij / NV, j = ij % NV; + ax2[i][j] = a[i][j] * x[j] * x[j]; + } + __syncthreads(); + + matmul_NT(ax2[0], a[0], ax2a[0]); + matmul_NT(ax2[0], c, ax2c); + cholesky_solve(ax2a, ax2c); + matmul_NN<1, NV, NC, SM_VER, BLOCK_DIM>(ax2c, a[0], r); + + if (tid < 32) { + d_max = 0.f; + for (int j = tid; j < NV; j += 32) { + float val = x[j] * (c[j] - r[j]); + d[j] = val; + d_max = fmaxf(d_max, val); + } + for (int offset = 16; offset > 0; offset >>= 1) { + d_max = fmaxf(d_max, __shfl_xor_sync(0xffffffff, d_max, offset)); + } + if (tid == 0) { + // Guard against d_max <= 0 from a degenerate / numerically-stuck + // iteration. A non-positive d_max would yield inf/NaN from the + // division and corrupt x on the next update. The 1.0 fallback + // produces a no-op step (x *= 1 - 1*0 = x) so the solver simply + // stalls rather than diverges, and the convergence check at the + // end of the kernel writes 0.5 if d_max stays small. + alpha = (d_max > 1e-9f) ? (0.999f / d_max) : 1.0f; + } + } + __syncthreads(); + + for (int j = tid; j < NV; j += dim) { + x[j] *= 1.f - alpha * d[j]; + } + __syncthreads(); + } + + // Compute residual ‖A x - b‖_inf for the convergence check. + matmul_NT(a[0], x, ax2c); + if (tid < 32) { + max_residual = 0.f; + for (int i = tid; i < NC; i += 32) { + max_residual = fmaxf(max_residual, fabsf(ax2c[i] - b[i])); + } + for (int offset = 16; offset > 0; offset >>= 1) { + max_residual = fmaxf(max_residual, __shfl_down_sync(0xffffffff, max_residual, offset)); + } + } + + auto& avail_flag = smem->avail_flag; + if (tid == 0) { + avail_flag = (d_max < 0.1f && x[NV - 1] >= 0.f && x[NV - 1] < 1e-4f && max_residual < 0.05f); + } + __syncthreads(); + + if (!avail_flag) { + for (int i = tid; i < NV; i += dim) { + result[i] = 0.5f; + } + } else { + for (int i = tid; i < NV; i += dim) { + result[i] = x[i]; + } + } +} + +template +void ipm_solve(tvm::ffi::TensorView A, tvm::ffi::TensorView b, tvm::ffi::TensorView c, tvm::ffi::TensorView result) { + using namespace host; + + SymbolicDevice device_; + TensorMatcher({NC, NV}).with_dtype().with_device(device_).verify(A); + TensorMatcher({NC}).with_dtype().with_device(device_).verify(b); + TensorMatcher({NV}).with_dtype().with_device(device_).verify(c); + TensorMatcher({NV}).with_dtype().with_device(device_).verify(result); + + const DLDevice device = device_.unwrap(); + const size_t smem_bytes = sizeof(ipm_smem); + + using KernelT = void (*)(float*, const float*, const float*, const float*); + KernelT kernel = ipm_solve_kernel; + + // Opt in to >48 KB dynamic shared memory if needed (Hopper supports up to + // 228 KB per block). + if (smem_bytes > 48 * 1024) { + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(smem_bytes)); + } + + LaunchKernel(/*grid_dim=*/1, /*block_dim=*/BLOCK_DIM, device, smem_bytes)( + kernel, + static_cast(result.data_ptr()), + static_cast(A.data_ptr()), + static_cast(b.data_ptr()), + static_cast(c.data_ptr())); +} + +} // namespace diff --git a/python/sglang/jit_kernel/csrc/lplb/lp_post.cuh b/python/sglang/jit_kernel/csrc/lplb/lp_post.cuh new file mode 100644 index 000000000..24e222da3 --- /dev/null +++ b/python/sglang/jit_kernel/csrc/lplb/lp_post.cuh @@ -0,0 +1,117 @@ +// LP post kernel: build the final log2phy_prob tensor from the IPM output +// `x` and the prep's `t1`. +// +// Python equivalent (~5 torch ops; this kernel is one launch): +// +// x_ratios = clamp(x[:NUM_RED_PHY], min=0) +// phy_prob = zeros(NUM_SINGLE + NUM_RED_PHY + 1) # +1 = sink slot +// phy_prob[phy_replicated] = x_ratios +// phy_prob[phy_single] = t1 +// log2phy_prob = take(phy_prob, log2phy) # (-1 wraps to sink) +// +// `log2phy` may contain -1 for unused replicas (DP-attention padding); we +// emulate torch.take's wrap-around by adding `phy_prob_size` to negative +// indices, which lands at the always-zero sink slot. +// +// Single-block launch. `phy_prob` lives in shared memory. + +#include +#include + +#include + +#include +#include + +#include + +namespace { + +template < + int NUM_LOGICAL, + int MAX_COPIES, + int NUM_SINGLE, + int NUM_RED_PHY, + int BLOCK_DIM> +__global__ void lp_post_kernel( + float* __restrict__ log2phy_prob, // (NUM_LOGICAL, MAX_COPIES) — written + const float* __restrict__ x, // (NV,) — IPM output + const float* __restrict__ t1, // (NUM_SINGLE,) — from prep + const int64_t* __restrict__ phy_single, // (NUM_SINGLE,) + const int64_t* __restrict__ phy_replicated, // (NUM_RED_PHY,) + const int64_t* __restrict__ log2phy) { // (NUM_LOGICAL, MAX_COPIES) + constexpr int PHY_PROB_SIZE = NUM_SINGLE + NUM_RED_PHY + 1; + + extern __shared__ unsigned char raw_smem[]; + float* phy_prob = reinterpret_cast(raw_smem); + + const int tid = threadIdx.x; + + // Stage 1: zero-init phy_prob (covers the sink slot at index PHY_PROB_SIZE-1). + for (int i = tid; i < PHY_PROB_SIZE; i += BLOCK_DIM) { + phy_prob[i] = 0.f; + } + __syncthreads(); + + // Stage 2: scatter x_ratios = clamp(x[:NUM_RED_PHY], min=0) at phy_replicated. + for (int i = tid; i < NUM_RED_PHY; i += BLOCK_DIM) { + int64_t idx = phy_replicated[i]; + phy_prob[idx] = fmaxf(x[i], 0.f); + } + // Stage 3: scatter t1 at phy_single. + for (int i = tid; i < NUM_SINGLE; i += BLOCK_DIM) { + int64_t idx = phy_single[i]; + phy_prob[idx] = t1[i]; + } + __syncthreads(); + + // Stage 4: gather log2phy_prob[i,j] = phy_prob[log2phy[i,j]]. + // -1 entries wrap to the sink slot (PHY_PROB_SIZE - 1), which is 0. + const int total = NUM_LOGICAL * MAX_COPIES; + for (int idx = tid; idx < total; idx += BLOCK_DIM) { + int64_t k = log2phy[idx]; + if (k < 0) k += PHY_PROB_SIZE; + log2phy_prob[idx] = phy_prob[k]; + } +} + +template +void lp_post( + tvm::ffi::TensorView log2phy_prob, + tvm::ffi::TensorView x, + tvm::ffi::TensorView t1, + tvm::ffi::TensorView phy_single, + tvm::ffi::TensorView phy_replicated, + tvm::ffi::TensorView log2phy) { + using namespace host; + + SymbolicDevice device_; + TensorMatcher({NUM_LOGICAL, MAX_COPIES}).with_dtype().with_device(device_).verify(log2phy_prob); + TensorMatcher({NUM_SINGLE}).with_dtype().with_device(device_).verify(t1); + TensorMatcher({NUM_SINGLE}).with_dtype().with_device(device_).verify(phy_single); + TensorMatcher({NUM_RED_PHY}).with_dtype().with_device(device_).verify(phy_replicated); + TensorMatcher({NUM_LOGICAL, MAX_COPIES}).with_dtype().with_device(device_).verify(log2phy); + // x has shape (NV,) which we don't constrain at this layer. + + constexpr int PHY_PROB_SIZE = NUM_SINGLE + NUM_RED_PHY + 1; + const size_t smem_bytes = PHY_PROB_SIZE * sizeof(float); + + using KernelT = void (*)(float*, const float*, const float*, const int64_t*, const int64_t*, const int64_t*); + KernelT kernel = lp_post_kernel; + + if (smem_bytes > 48 * 1024) { + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(smem_bytes)); + } + + const DLDevice device = device_.unwrap(); + LaunchKernel(/*grid=*/1, /*block=*/BLOCK_DIM, device, smem_bytes)( + kernel, + static_cast(log2phy_prob.data_ptr()), + static_cast(x.data_ptr()), + static_cast(t1.data_ptr()), + static_cast(phy_single.data_ptr()), + static_cast(phy_replicated.data_ptr()), + static_cast(log2phy.data_ptr())); +} + +} // namespace diff --git a/python/sglang/jit_kernel/csrc/lplb/lp_prep.cuh b/python/sglang/jit_kernel/csrc/lplb/lp_prep.cuh new file mode 100644 index 000000000..6cd5e8b79 --- /dev/null +++ b/python/sglang/jit_kernel/csrc/lplb/lp_prep.cuh @@ -0,0 +1,182 @@ +// LP prep kernel: build IPM inputs (A_full last column, b, t1) from +// global_counts and the constant per-rank metadata (log_single, +// log_replicated, B1, A_base_row_sum). +// +// Python equivalent (~8 torch ops; this kernel is one launch): +// +// total = global_counts.sum() +// counts_norm = global_counts / total.clamp(min=1.0) +// t1 = counts_norm[log_single] # (NUM_SINGLE,) +// b1 = counts_norm[log_replicated] # (NUM_RED_LOG,) +// b2 = -(B1 @ t1).flatten() # (NUM_GPUS,) +// b = cat(b1, b2) # (NC,) +// A_full[:, -1] = b - A_base_row_sum # last column only +// +// `A_full` is pre-allocated by the caller with shape (NC, NV); its first +// NV-1 columns are pre-filled with A_base.copy_() at solver init and not +// touched by this kernel. +// +// Single-block launch. All intermediate state lives in shared memory. + +#include +#include + +#include + +#include +#include + +#include + +namespace { + +template +__global__ void lp_prep_kernel( + float* __restrict__ A_full, // (NC, NV) — last column is written + float* __restrict__ b, // (NC,) — written + float* __restrict__ t1, // (NUM_SINGLE,) — written + const float* __restrict__ global_counts, // (num_logical,) + const int64_t* __restrict__ log_single, // (NUM_SINGLE,) + const int64_t* __restrict__ log_replicated, // (NUM_RED_LOG,) + const float* __restrict__ B1, // (NUM_GPUS, NUM_SINGLE) + const float* __restrict__ A_base_row_sum) { // (NC,) + static_assert(NC == NUM_RED_LOG + NUM_GPUS, "NC must equal NUM_RED_LOG + NUM_GPUS"); + constexpr int WARPS_PER_BLOCK = BLOCK_DIM / 32; + + extern __shared__ unsigned char raw_smem[]; + // Layout: shared_t1 [NUM_SINGLE] | shared_b1 [NUM_RED_LOG] | reduce_buf [WARPS_PER_BLOCK] | total_pad [1] + float* shared_t1 = reinterpret_cast(raw_smem); + float* shared_b1 = shared_t1 + NUM_SINGLE; + float* reduce_buf = shared_b1 + NUM_RED_LOG; + float* shared_total = reduce_buf + WARPS_PER_BLOCK; + + const int tid = threadIdx.x; + const int lane = tid & 31; + const int warp_id = tid >> 5; + + // ---- Stage 1: gather raw t1 / b1 + partial sum for total ---- + float local_sum = 0.f; + for (int i = tid; i < NUM_SINGLE; i += BLOCK_DIM) { + float v = global_counts[log_single[i]]; + shared_t1[i] = v; // raw, scaled below + local_sum += v; + } + for (int i = tid; i < NUM_RED_LOG; i += BLOCK_DIM) { + float v = global_counts[log_replicated[i]]; + shared_b1[i] = v; + local_sum += v; + } + + // Block-level reduction: warp shuffle -> shared mem -> warp 0 final reduce. + for (int offset = 16; offset > 0; offset >>= 1) { + local_sum += __shfl_xor_sync(0xffffffff, local_sum, offset); + } + if (lane == 0) { + reduce_buf[warp_id] = local_sum; + } + __syncthreads(); + if (warp_id == 0) { + float v = (tid < WARPS_PER_BLOCK) ? reduce_buf[tid] : 0.f; + for (int offset = 16; offset > 0; offset >>= 1) { + v += __shfl_xor_sync(0xffffffff, v, offset); + } + if (tid == 0) { + shared_total[0] = fmaxf(v, 1.0f); // clamp(min=1.0) + } + } + __syncthreads(); + const float inv_total = 1.0f / shared_total[0]; + + // ---- Stage 2: scale t1 (keep in shmem for matmul, also write to global) + // scale b[0..NUM_RED_LOG] = scaled b1 + for (int i = tid; i < NUM_SINGLE; i += BLOCK_DIM) { + float v = shared_t1[i] * inv_total; + shared_t1[i] = v; + t1[i] = v; + } + for (int i = tid; i < NUM_RED_LOG; i += BLOCK_DIM) { + b[i] = shared_b1[i] * inv_total; + } + __syncthreads(); + + // ---- Stage 3: b[NUM_RED_LOG + j] = -(B1[j] · t1) for j in [0, NUM_GPUS). + // Sequential GEMV across NUM_GPUS=16 outputs; each is a 240-wide dot + // product reduced across the block. Cheap at this size. + for (int j = 0; j < NUM_GPUS; j++) { + float dot = 0.f; + for (int k = tid; k < NUM_SINGLE; k += BLOCK_DIM) { + dot += B1[j * NUM_SINGLE + k] * shared_t1[k]; + } + for (int offset = 16; offset > 0; offset >>= 1) { + dot += __shfl_xor_sync(0xffffffff, dot, offset); + } + if (lane == 0) { + reduce_buf[warp_id] = dot; + } + __syncthreads(); + if (warp_id == 0) { + float v = (tid < WARPS_PER_BLOCK) ? reduce_buf[tid] : 0.f; + for (int offset = 16; offset > 0; offset >>= 1) { + v += __shfl_xor_sync(0xffffffff, v, offset); + } + if (tid == 0) { + b[NUM_RED_LOG + j] = -v; + } + } + __syncthreads(); + } + + // ---- Stage 4: A_full[i][NV-1] = b[i] - A_base_row_sum[i]. + // First NV-1 columns of A_full are pre-filled with A_base at solver init, + // so we only write the last column here. + for (int i = tid; i < NC; i += BLOCK_DIM) { + A_full[i * NV + (NV - 1)] = b[i] - A_base_row_sum[i]; + } +} + +template +void lp_prep( + tvm::ffi::TensorView A_full, + tvm::ffi::TensorView b, + tvm::ffi::TensorView t1, + tvm::ffi::TensorView global_counts, + tvm::ffi::TensorView log_single, + tvm::ffi::TensorView log_replicated, + tvm::ffi::TensorView B1, + tvm::ffi::TensorView A_base_row_sum) { + using namespace host; + + SymbolicDevice device_; + TensorMatcher({NC, NV}).with_dtype().with_device(device_).verify(A_full); + TensorMatcher({NC}).with_dtype().with_device(device_).verify(b); + TensorMatcher({NUM_SINGLE}).with_dtype().with_device(device_).verify(t1); + TensorMatcher({NUM_SINGLE}).with_dtype().with_device(device_).verify(log_single); + TensorMatcher({NUM_RED_LOG}).with_dtype().with_device(device_).verify(log_replicated); + TensorMatcher({NUM_GPUS, NUM_SINGLE}).with_dtype().with_device(device_).verify(B1); + TensorMatcher({NC}).with_dtype().with_device(device_).verify(A_base_row_sum); + + constexpr int WARPS_PER_BLOCK = BLOCK_DIM / 32; + const size_t smem_bytes = (NUM_SINGLE + NUM_RED_LOG + WARPS_PER_BLOCK + 1) * sizeof(float); + + using KernelT = + void (*)(float*, float*, float*, const float*, const int64_t*, const int64_t*, const float*, const float*); + KernelT kernel = lp_prep_kernel; + + if (smem_bytes > 48 * 1024) { + cudaFuncSetAttribute(kernel, cudaFuncAttributeMaxDynamicSharedMemorySize, static_cast(smem_bytes)); + } + + const DLDevice device = device_.unwrap(); + LaunchKernel(/*grid=*/1, /*block=*/BLOCK_DIM, device, smem_bytes)( + kernel, + static_cast(A_full.data_ptr()), + static_cast(b.data_ptr()), + static_cast(t1.data_ptr()), + static_cast(global_counts.data_ptr()), + static_cast(log_single.data_ptr()), + static_cast(log_replicated.data_ptr()), + static_cast(B1.data_ptr()), + static_cast(A_base_row_sum.data_ptr())); +} + +} // namespace diff --git a/python/sglang/jit_kernel/lplb/__init__.py b/python/sglang/jit_kernel/lplb/__init__.py new file mode 100644 index 000000000..2c6f32ed9 --- /dev/null +++ b/python/sglang/jit_kernel/lplb/__init__.py @@ -0,0 +1 @@ +"""LPLB (LP-based Load Balancer) JIT kernels for expert parallelism.""" diff --git a/python/sglang/jit_kernel/lplb/cublasdx_solver.py b/python/sglang/jit_kernel/lplb/cublasdx_solver.py new file mode 100644 index 000000000..889fecd20 --- /dev/null +++ b/python/sglang/jit_kernel/lplb/cublasdx_solver.py @@ -0,0 +1,14 @@ +"""Backwards-compatible shim. + +The Numba/nvmath-python fused IPM that used to live here has been replaced +by a CUDA C++ kernel JIT-compiled via sglang's ``load_jit`` infrastructure. +The new implementation lives in ``cuda_solver``. This module re-exports the +public API so any external import keeps working. +""" + +from sglang.jit_kernel.lplb.cuda_solver import ( # noqa: F401 + solve_ipm, + warmup, +) + +__all__ = ["solve_ipm", "warmup"] diff --git a/python/sglang/jit_kernel/lplb/cuda_solver.py b/python/sglang/jit_kernel/lplb/cuda_solver.py new file mode 100644 index 000000000..60a78c751 --- /dev/null +++ b/python/sglang/jit_kernel/lplb/cuda_solver.py @@ -0,0 +1,324 @@ +"""JIT-compiled CUDA Interior Point Method LP solver. + +Replaces the Numba/nvmath-python implementation in ``cublasdx_solver.py``. +The kernel is a single-block fused IPM defined in +``csrc/lplb/ipm.cuh`` and compiled per ``(NC, NV, BLOCK_DIM, SM_VER, +NUM_ITERS)`` tuple via sglang's ``tvm-ffi`` ``load_jit``. + +Per-call CPU overhead is dominated by the pybind11 dispatch + four +``data_ptr()`` calls (~5–10 µs total), versus ~500–700 µs for the prior +Numba path (numba dispatcher chain + ``as_cuda_array`` per tensor). +""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.jit_kernel.utils import ( + cache_once, + get_jit_cuda_arch, + load_jit, + make_cpp_args, +) + +if TYPE_CHECKING: + from tvm_ffi.module import Module + +logger = logging.getLogger(__name__) + +DEFAULT_BLOCK_DIM = 256 +# Per-element kernels (post-LP dispatch) saturate easily — also 256. +DISPATCH_BLOCK_DIM = 256 +DEFAULT_NUM_ITERS = 5 + + +def _sm_ver() -> int: + arch = get_jit_cuda_arch() + return arch.major * 100 + arch.minor * 10 + + +@cache_once +def _ipm_module( + nc: int, nv: int, block_dim: int, num_iters: int, sm_ver: int +) -> Module: + """JIT-compile the IPM kernel for one shape. Cached for the process lifetime.""" + args = make_cpp_args(nc, nv, block_dim, sm_ver, num_iters) + # The kernel uses cuBLASDx (header-only) for the GEMMs and a hand-written + # block-level Cholesky for the POSV. No -rdc=true / static-lib linkage + # required, so sglang's tvm-ffi load_jit handles the build with the + # default flags. + return load_jit( + "lplb_ipm", + *args, + cuda_files=["lplb/ipm.cuh"], + cuda_wrappers=[("ipm_solve", f"ipm_solve<{args}>")], + extra_dependencies=["mathdx"], + ) + + +def warmup( + nc: int, + nv: int, + num_iters: int = DEFAULT_NUM_ITERS, + device: str = "cuda", +) -> None: + """JIT-compile the kernel for ``(nc, nv)`` so the first real solve isn't + paying the compile cost. Raises on compile or launch failure. + """ + module = _ipm_module(nc, nv, DEFAULT_BLOCK_DIM, num_iters, _sm_ver()) + # Trigger any first-call lazy initialization. + A = torch.zeros(nc, nv, dtype=torch.float32, device=device) + b = torch.zeros(nc, dtype=torch.float32, device=device) + c = torch.zeros(nv, dtype=torch.float32, device=device) + result = torch.empty(nv, dtype=torch.float32, device=device) + module.ipm_solve(A, b, c, result) + logger.info(f"LPLB CUDA IPM solver: warmed up for (NC={nc}, NV={nv})") + + +def solve_ipm( + A: torch.Tensor, + b: torch.Tensor, + c: torch.Tensor, + num_iters: int = DEFAULT_NUM_ITERS, + result: torch.Tensor | None = None, +) -> torch.Tensor: + """Run the fused single-SM IPM kernel. + + cuBLASDx GEMMs + hand-written block Cholesky, dispatched per the + module docstring. + + Args: + A: Constraint matrix, shape ``(NC, NV)``, float32, on CUDA. + b: RHS vector, shape ``(NC,)``, float32, on CUDA. + c: Objective coefficients, shape ``(NV,)``, float32, on CUDA. + num_iters: Number of barrier iterations (default 5). + result: Optional pre-allocated ``(NV,)`` float32 CUDA buffer to write + into. When omitted the kernel allocates a fresh result tensor + (~20 µs of CPU overhead). Passing in a long-lived buffer skips + that alloc on every solve. + + Returns: + x: Solution vector, shape ``(NV,)``, float32. The kernel writes 0.5 + for every entry on non-convergence (matches the prior Numba behavior). + """ + assert A.is_cuda and b.is_cuda and c.is_cuda + assert A.dtype == torch.float32 + nc, nv = A.shape + assert b.shape == (nc,), f"b shape mismatch: {b.shape} vs ({nc},)" + assert c.shape == (nv,), f"c shape mismatch: {c.shape} vs ({nv},)" + + module = _ipm_module(nc, nv, DEFAULT_BLOCK_DIM, num_iters, _sm_ver()) + if result is None: + result = torch.empty(nv, dtype=torch.float32, device=A.device) + module.ipm_solve(A, b, c, result) + return result + + +@cache_once +def _prep_module( + nc: int, + nv: int, + num_single: int, + num_red_log: int, + num_gpus: int, + block_dim: int, +) -> Module: + args = make_cpp_args(nc, nv, num_single, num_red_log, num_gpus, block_dim) + return load_jit( + "lplb_lp_prep", + *args, + cuda_files=["lplb/lp_prep.cuh"], + cuda_wrappers=[("lp_prep", f"lp_prep<{args}>")], + ) + + +def prep_lp_inputs( + A_full: torch.Tensor, + b: torch.Tensor, + t1: torch.Tensor, + global_counts: torch.Tensor, + log_single: torch.Tensor, + log_replicated: torch.Tensor, + B1: torch.Tensor, + A_base_row_sum: torch.Tensor, +) -> None: + """Replace the 8 torch ops that built the IPM inputs with one CUDA kernel. + + Writes into the caller-provided ``A_full`` (last column only), ``b``, + and ``t1`` buffers. ``A_full``'s first ``NV-1`` columns must already + hold ``A_base.copy_()`` from solver init — this kernel does not touch + them. + """ + nc, nv = A_full.shape + num_single = log_single.shape[0] + num_red_log = log_replicated.shape[0] + num_gpus, _ns = B1.shape + module = _prep_module(nc, nv, num_single, num_red_log, num_gpus, DEFAULT_BLOCK_DIM) + module.lp_prep( + A_full, b, t1, global_counts, log_single, log_replicated, B1, A_base_row_sum + ) + + +@cache_once +def _post_module( + num_logical: int, + max_copies: int, + num_single: int, + num_red_phy: int, + block_dim: int, +) -> Module: + args = make_cpp_args(num_logical, max_copies, num_single, num_red_phy, block_dim) + return load_jit( + "lplb_lp_post", + *args, + cuda_files=["lplb/lp_post.cuh"], + cuda_wrappers=[("lp_post", f"lp_post<{args}>")], + ) + + +def extract_log2phy_prob( + log2phy_prob: torch.Tensor, + x: torch.Tensor, + t1: torch.Tensor, + phy_single: torch.Tensor, + phy_replicated: torch.Tensor, + log2phy: torch.Tensor, +) -> None: + """Replace the 5 torch ops that turned the IPM output into log2phy_prob + with one CUDA kernel. Writes into the caller-provided ``log2phy_prob`` + buffer of shape ``(num_logical, max_copies)``. + """ + num_logical, max_copies = log2phy_prob.shape + num_single = phy_single.shape[0] + num_red_phy = phy_replicated.shape[0] + module = _post_module( + num_logical, max_copies, num_single, num_red_phy, DEFAULT_BLOCK_DIM + ) + module.lp_post(log2phy_prob, x, t1, phy_single, phy_replicated, log2phy) + + +@cache_once +def _dispatch_module(max_copies: int, block_dim: int) -> Module: + args = make_cpp_args(max_copies, block_dim) + return load_jit( + "lplb_dispatch_probability", + *args, + cuda_files=["lplb/dispatch_probability.cuh"], + cuda_wrappers=[("dispatch_probability", f"dispatch_probability<{args}>")], + ) + + +def dispatch_probability( + topk_ids: torch.Tensor, + log2phy_prob: torch.Tensor, + log2phy_map: torch.Tensor, + random_vals: Optional[torch.Tensor] = None, +) -> torch.Tensor: + """Replace the 7 torch ops in `_topk_ids_logical_to_physical_probability` + with a single per-token-per-slot CUDA kernel. + + Samples a physical expert per (token, slot) via inverse-CDF on the + per-row LP probabilities. Bit-equivalent to + :func:`dispatch_probability_torch_reference` when given the same + ``random_vals`` (modulo float rounding in the cumulative sum). + + Args: + topk_ids: (num_tokens, topk) int32, on CUDA. Logical expert ids from + the router. + log2phy_prob: (num_logical, max_copies) float32. LP solver output. + log2phy_map: (num_logical, max_copies) int32. -1 entries are + unused replicas; treated as 0-weight in the multinomial. + random_vals: Optional (N,) float32 CUDA tensor of uniform samples in + [0, 1). When omitted, the function generates fresh values via + ``torch.rand``. Pass explicitly when comparing against the torch + reference for deterministic equivalence. + + Returns: + Physical topk ids tensor with the same shape as ``topk_ids``. + """ + original_shape = topk_ids.shape + flat_ids = topk_ids.reshape(-1).contiguous().to(torch.int32) + n = flat_ids.shape[0] + num_logical, max_copies = log2phy_prob.shape + assert log2phy_map.shape == (num_logical, max_copies) + map32 = log2phy_map.contiguous().to(torch.int32) + + out = torch.empty(n, dtype=torch.int32, device=topk_ids.device) + if random_vals is None: + random_vals = torch.rand(n, dtype=torch.float32, device=topk_ids.device) + else: + assert random_vals.shape == ( + n, + ), f"random_vals must be shape ({n},), got {tuple(random_vals.shape)}" + module = _dispatch_module(max_copies, DISPATCH_BLOCK_DIM) + module.dispatch_probability(out, flat_ids, log2phy_prob, map32, random_vals) + return out.view(original_shape).to(topk_ids.dtype) + + +def dispatch_probability_torch_reference( + topk_ids: torch.Tensor, + log2phy_prob: torch.Tensor, + log2phy_map: torch.Tensor, + random_vals: torch.Tensor, +) -> torch.Tensor: + """Pure-torch reference of :func:`dispatch_probability`. + + Mirrors the CUDA kernel's algorithm exactly (inverse-CDF via cumsum + + threshold) so the two paths are bit-equivalent for identical + ``random_vals``, modulo floating-point rounding in the cumsum. Kept + for numerical comparison and testing — not on the production hot + path (it allocates and runs ~8 torch ops; the fused kernel collapses + them into one launch). + + Algorithm (matches ``csrc/lplb/dispatch_probability.cuh``): + + 1. Gather the per-row probability vector and physical-id map for + each logical id in ``topk_ids``. + 2. If the row sum is zero (LP gave no signal), fall back to + uniform over valid replicas (``log2phy_map != -1``). + 3. Sample: smallest ``c`` such that ``cumsum[0..c] > u * row_sum``, + where ``u = random_vals[i]``. Ties favor advancing ``c``, + matching the CUDA kernel. + 4. Return ``log2phy_map[logical_id, c]``. + + Args: + topk_ids: (num_tokens, topk) int, on CUDA or CPU. Logical expert ids. + log2phy_prob: (num_logical, max_copies) float32. LP solver output. + log2phy_map: (num_logical, max_copies) int. -1 = unused replica. + random_vals: (N,) float32, where N = ``topk_ids.numel()``. Uniform + samples in [0, 1) — same shape and semantics as the CUDA kernel. + + Returns: + Physical topk ids tensor with the same shape and dtype as ``topk_ids``. + """ + original_shape = topk_ids.shape + flat_ids = topk_ids.reshape(-1).long() + n = flat_ids.shape[0] + num_logical, max_copies = log2phy_prob.shape + assert log2phy_map.shape == (num_logical, max_copies) + assert random_vals.shape == ( + n, + ), f"random_vals must be shape ({n},), got {tuple(random_vals.shape)}" + + # Gather per-row probabilities and physical maps. + probs = log2phy_prob[flat_ids] # (N, max_copies), float32 + maps = log2phy_map[flat_ids] # (N, max_copies), same int dtype as input + + # Fallback when row_sum == 0: uniform over valid replicas. + row_sum = probs.sum(dim=-1, keepdim=True) # (N, 1) + fallback_probs = (maps >= 0).to(probs.dtype) # (N, max_copies) + probs = torch.where(row_sum > 0, probs, fallback_probs) + row_sum = probs.sum(dim=-1) # (N,) + + # Inverse-CDF sample: smallest c such that cumsum[..c] > u. + # ``(cum <= u).sum(dim=-1)`` counts how many slots are still below u, + # which equals the CUDA kernel's ``chosen`` after its for-loop. + u = (random_vals * row_sum).unsqueeze(-1) # (N, 1) + cum = probs.cumsum(dim=-1) # (N, max_copies) + chosen = (cum <= u).sum(dim=-1).clamp(max=max_copies - 1) # (N,) + + out = maps.gather(1, chosen.unsqueeze(-1)).squeeze(-1) + return out.view(original_shape).to(topk_ids.dtype) diff --git a/python/sglang/jit_kernel/lplb/shmem_budget.py b/python/sglang/jit_kernel/lplb/shmem_budget.py new file mode 100644 index 000000000..c91f04e19 --- /dev/null +++ b/python/sglang/jit_kernel/lplb/shmem_budget.py @@ -0,0 +1,152 @@ +"""Shared-memory budget accounting for the fused IPM kernel. + +Fused layout (fp32), one block per LP, all state in shared memory:: + + A NC * NV constraint matrix (resident) + c NV cost vector (resident) + x NV IPM state (resident) + ata NC * NC KKT matrix / Cholesky factor + rhs NC ax2c, then delta + d NV aliased with r = A.T @ delta + + S_elems = NC*NV + NC*NC + 3*NV + NC + +Dynamic shared-memory cap per block (with opt-in via +``cudaFuncAttributeMaxDynamicSharedMemorySize``): + + A100 SM_80 164 KB practical 160 KB + H100 SM_90 227 KB practical 223 KB <- default target + H200 SM_90 227 KB + H20 SM_90 227 KB + B200 SM_100 228 KB practical 224 KB +""" + +from __future__ import annotations + +from dataclasses import dataclass + +# Per-block slack reserved for cuBLASDx workspace and CUDA runtime state. +_RUNTIME_PAD_BYTES = 256 + +# fp32 +_BYTES_PER_ELEM = 4 + +# Practical per-block dynamic shmem caps (bytes) +GPU_BUDGETS_BYTES: dict[str, int] = { + "a100": 160 * 1024, + "h100": 223 * 1024, + "h200": 223 * 1024, + "h20": 223 * 1024, + "b200": 224 * 1024, +} + + +@dataclass(frozen=True) +class ShmemBreakdown: + nc: int + nv: int + a_bytes: int + c_bytes: int + x_bytes: int + ata_bytes: int + rhs_bytes: int + d_bytes: int + pad_bytes: int + + @property + def total_bytes(self) -> int: + return ( + self.a_bytes + + self.c_bytes + + self.x_bytes + + self.ata_bytes + + self.rhs_bytes + + self.d_bytes + + self.pad_bytes + ) + + def as_kib(self) -> float: + return self.total_bytes / 1024.0 + + +def shmem_bytes(nc: int, nv: int, bytes_per_elem: int = _BYTES_PER_ELEM) -> int: + """Exact byte count for the fused layout with the given (NC, NV).""" + return bytes_per_elem * (nc * nv + nc * nc + 3 * nv + nc) + _RUNTIME_PAD_BYTES + + +def breakdown( + nc: int, nv: int, bytes_per_elem: int = _BYTES_PER_ELEM +) -> ShmemBreakdown: + """Per-array byte breakdown — useful for debugging shmem pressure.""" + b = bytes_per_elem + return ShmemBreakdown( + nc=nc, + nv=nv, + a_bytes=b * nc * nv, + c_bytes=b * nv, + x_bytes=b * nv, + ata_bytes=b * nc * nc, + rhs_bytes=b * nc, + d_bytes=b * nv, + pad_bytes=_RUNTIME_PAD_BYTES, + ) + + +def gpu_budget_bytes(gpu: str) -> int: + key = gpu.lower() + if key not in GPU_BUDGETS_BYTES: + raise ValueError( + f"unknown gpu '{gpu}', expected one of {sorted(GPU_BUDGETS_BYTES)}" + ) + return GPU_BUDGETS_BYTES[key] + + +def fits(nc: int, nv: int, gpu: str = "h100") -> bool: + return shmem_bytes(nc, nv) <= gpu_budget_bytes(gpu) + + +def assert_fits(nc: int, nv: int, gpu: str = "h100") -> None: + """Raise if the fused kernel will not fit on the target GPU.""" + used = shmem_bytes(nc, nv) + cap = gpu_budget_bytes(gpu) + if used > cap: + raise ValueError( + f"fused IPM kernel needs {used/1024:.1f} KiB of shared memory for " + f"NC={nc}, NV={nv}, but {gpu} allows {cap/1024:.1f} KiB/block. " + f"Either reduce problem size or switch to a tiled design." + ) + + +def max_nc_for_nv(nv: int, gpu: str = "h100") -> int: + """Largest NC that fits for a given NV. Solves + 4 * (NC^2 + (NV+1)*NC + 3*NV) + pad <= cap + via the quadratic formula (monotone in NC). Returns 0 if even NC=1 overflows. + """ + cap = gpu_budget_bytes(gpu) + b = _BYTES_PER_ELEM + # cap - pad >= b * (NC^2 + (NV+1)*NC + 3*NV) + rhs = (cap - _RUNTIME_PAD_BYTES) / b - 3 * nv + if rhs <= 0: + return 0 + # NC^2 + (NV+1)*NC - rhs <= 0 + import math + + disc = (nv + 1) ** 2 + 4 * rhs + nc_max = int((-(nv + 1) + math.sqrt(disc)) / 2.0) + while nc_max > 0 and shmem_bytes(nc_max, nv) > cap: + nc_max -= 1 + return max(nc_max, 0) + + +def report(nc: int, nv: int, gpu: str = "h100") -> str: + """Human-readable summary — used by kernels on init for logging.""" + bd = breakdown(nc, nv) + cap = gpu_budget_bytes(gpu) + status = "FITS" if bd.total_bytes <= cap else "OVER BUDGET" + return ( + f"[shmem] NC={nc} NV={nv} gpu={gpu} | " + f"A={bd.a_bytes/1024:.1f}K " + f"ata={bd.ata_bytes/1024:.1f}K " + f"rest={(bd.c_bytes+bd.x_bytes+bd.rhs_bytes+bd.d_bytes)/1024:.1f}K | " + f"total={bd.total_bytes/1024:.1f}K / {cap/1024:.1f}K {status}" + ) diff --git a/python/sglang/jit_kernel/lplb/torch_solver.py b/python/sglang/jit_kernel/lplb/torch_solver.py new file mode 100644 index 000000000..337151b34 --- /dev/null +++ b/python/sglang/jit_kernel/lplb/torch_solver.py @@ -0,0 +1,192 @@ +"""IPM LP Solver entry point — dispatches to the fused JIT CUDA kernel. + +Solves: min c^T x subject to Ax = b, x >= 0 +using a barrier (interior point) method with 5 iterations. + +The fused kernel lives in ``cuda_solver`` (CUDA C++ via ``load_jit``, +backed by header-only cuBLASDx + a hand-written block Cholesky). This +module is the public-facing import surface for callers (``LPLBSolver``) +and resolves/caches the backend on first use. + +LPLB requires Hopper-class hardware and Math-DX cuBLASDx headers. If +either is missing, ``warmup`` and ``solve_ipm`` raise — there is no +silent fallback. +""" + +import logging + +import torch + +logger = logging.getLogger(__name__) + + +# Backend dispatch state (resolved on first call, cached afterwards) +_BACKEND_CHECKED = False +_FUSED_AVAILABLE = False +_FUSED_SOLVE_IPM = None # type: ignore[assignment] +_FUSED_WARMUP = None # type: ignore[assignment] +_FUSED_ASSERT_FITS = None # type: ignore[assignment] + + +def _init_fused_backend() -> None: + """Resolve the fused backend once. Records WHY it's disabled when it is.""" + global _BACKEND_CHECKED, _FUSED_AVAILABLE + global _FUSED_SOLVE_IPM, _FUSED_WARMUP, _FUSED_ASSERT_FITS + + if _BACKEND_CHECKED: + return + _BACKEND_CHECKED = True + + if not torch.cuda.is_available(): + logger.info("LPLB fused solver disabled: CUDA not available") + return + + cap = torch.cuda.get_device_capability() + if cap[0] < 9: + logger.info( + f"LPLB fused solver disabled: GPU SM {cap[0]}.{cap[1]} < 9.0 " + "(requires Hopper or newer)" + ) + return + + try: + from sglang.jit_kernel.lplb.cuda_solver import solve_ipm as fused_solve_ipm + from sglang.jit_kernel.lplb.cuda_solver import warmup as fused_warmup + from sglang.jit_kernel.lplb.shmem_budget import assert_fits + except ImportError as e: + logger.info( + f"LPLB fused solver disabled: {e}. " + "Install Math-DX cuBLASDx via `pip install nvidia-mathdx` " + "or set MATHDX_HOME to an extracted archive." + ) + return + + _FUSED_SOLVE_IPM = fused_solve_ipm + _FUSED_WARMUP = fused_warmup + _FUSED_ASSERT_FITS = assert_fits + _FUSED_AVAILABLE = True + logger.info("LPLB fused solver enabled (CUDA C++ via load_jit, cuBLASDx)") + + +def _unavailable_reason() -> str: + if not torch.cuda.is_available(): + return "CUDA is not available" + cap = torch.cuda.get_device_capability() + if cap[0] < 9: + return f"GPU SM {cap[0]}.{cap[1]} < 9.0 (requires Hopper or newer)" + return ( + "Math-DX cuBLASDx headers not found — install via " + "`pip install nvidia-mathdx` or set MATHDX_HOME" + ) + + +def warmup(nc: int, nv: int, num_iters: int = 5, device: str = "cuda") -> None: + """Pre-JIT-compile the fused kernel for a given (NC, NV) shape. + + Call once per unique shape at solver construction time to hide the + 20-40s JIT compilation cost. Raises if the fused backend is + unavailable, the shape exceeds the shmem budget, or the kernel + fails to compile/launch. + """ + _init_fused_backend() + if not _FUSED_AVAILABLE: + raise RuntimeError(f"LPLB fused solver unavailable: {_unavailable_reason()}") + _FUSED_ASSERT_FITS(nc, nv, gpu="h100") + _FUSED_WARMUP(nc, nv, num_iters=num_iters, device=device) + + +def solve_ipm( + A: torch.Tensor, + b: torch.Tensor, + c: torch.Tensor, + num_iters: int = 5, +) -> torch.Tensor: + """Barrier-method Interior Point solver for standard-form LP. + + Dispatches to the JIT-compiled CUDA C++ kernel (Hopper+ GPU with + Math-DX cuBLASDx headers, reachable via ``nvidia-mathdx`` PyPI + package or ``MATHDX_HOME``). Raises if the fused backend is + unavailable or the inputs aren't on CUDA in float32. + + Args: + A: Constraint matrix, shape (NC, NV), float32, on CUDA. + b: RHS vector, shape (NC,), float32, on CUDA. + c: Objective coefficients, shape (NV,), float32, on CUDA. + num_iters: Number of barrier iterations (default 5). + + Returns: + x: Solution vector, shape (NV,), float32. The kernel writes 0.5 + for every entry on non-convergence. + """ + nc, nv = A.shape + assert b.shape == (nc,), f"b shape mismatch: {b.shape} vs ({nc},)" + assert c.shape == (nv,), f"c shape mismatch: {c.shape} vs ({nv},)" + + _init_fused_backend() + if not _FUSED_AVAILABLE: + raise RuntimeError(f"LPLB fused solver unavailable: {_unavailable_reason()}") + if not A.is_cuda: + raise RuntimeError( + f"LPLB fused solver requires CUDA tensors; got A on {A.device}." + ) + if A.dtype != torch.float32: + raise RuntimeError( + f"LPLB fused solver requires float32; got A.dtype={A.dtype}." + ) + return _FUSED_SOLVE_IPM(A, b, c, num_iters=num_iters) + + +def solve_ipm_torch_reference( + A: torch.Tensor, + b: torch.Tensor, + c: torch.Tensor, + num_iters: int = 5, +) -> torch.Tensor: + """Pure-torch reference for the fused IPM kernel — testing only. + + Mirrors the barrier-method iteration in ``csrc/lplb/ipm.cuh`` + step-for-step so the two can be compared numerically: + + x <- 1 + for _ in range(num_iters): + ax2 = A * x^2 # (NC, NV) + ax2a = ax2 @ A^T # (NC, NC) KKT matrix + delta = solve(ax2a, ax2 @ c) + r = delta^T @ A # (NV,) + d = x * (c - r) + alpha = 0.999 / d_max (or 1.0 if d_max <= 1e-9) + x *= 1 - alpha * d + write 0.5 everywhere on non-convergence. + + NOT bit-equivalent to the kernel: the kernel factors the KKT system + with a hand-written block Cholesky while this uses + ``torch.linalg.solve`` (LU). The two agree to a small tolerance + (the numerical difference being the whole point of the comparison + test). This function is never on the production path — the fused + kernel is the only LP solver at runtime. + """ + nc, nv = A.shape + assert b.shape == (nc,), f"b shape mismatch: {b.shape} vs ({nc},)" + assert c.shape == (nv,), f"c shape mismatch: {c.shape} vs ({nv},)" + + x = torch.ones(nv, device=A.device, dtype=torch.float32) + d_max = torch.tensor(0.0, device=A.device, dtype=torch.float32) + for _ in range(num_iters): + ax2 = A * (x * x).unsqueeze(0) # (NC, NV) + ax2a = ax2 @ A.t() # (NC, NC) + ax2c = ax2 @ c # (NC,) + # Match the kernel's 1e-12 pivot clamp via a tiny diagonal jitter so + # a (near-)singular KKT system stays solvable instead of raising. + ax2a = ax2a + 1e-12 * torch.eye(nc, device=A.device, dtype=torch.float32) + delta = torch.linalg.solve(ax2a, ax2c) # (NC,) + r = A.t() @ delta # (NV,) + d = x * (c - r) # (NV,) + d_max = d.max() + alpha = 0.999 / d_max if d_max > 1e-9 else torch.tensor(1.0, device=A.device) + x = x * (1.0 - alpha * d) + + max_residual = (A @ x - b).abs().max() + converged = (d_max < 0.1) and (0 <= x[-1] < 1e-4) and (max_residual < 0.05) + if not converged: + return torch.full((nv,), 0.5, device=A.device, dtype=torch.float32) + return x diff --git a/python/sglang/jit_kernel/utils.py b/python/sglang/jit_kernel/utils.py index 2f5178f56..bdff4bb67 100644 --- a/python/sglang/jit_kernel/utils.py +++ b/python/sglang/jit_kernel/utils.py @@ -472,6 +472,53 @@ def get_flashinfer_include_paths() -> List[str]: return include_paths +def get_mathdx_root() -> Optional[pathlib.Path]: + """Locate the NVIDIA Math-DX install (cuBLASDx headers). + + Searches in order: + 1. ``$MATHDX_HOME`` env var (extracted Math-DX archive root). + 2. The ``nvidia-mathdx`` PyPI package, if installed. + """ + env_home = os.environ.get("MATHDX_HOME") + if env_home: + candidate = pathlib.Path(env_home).expanduser().resolve() + if (candidate / "include").exists(): + return candidate + + # The ``nvidia-mathdx`` wheel installs as the namespace package + # ``nvidia.mathdx`` (no __init__, so spec.origin is None); resolve it via + # submodule_search_locations rather than _find_package_root, which only + # handles regular packages. + spec = importlib.util.find_spec("nvidia.mathdx") + if spec is not None: + roots = list(spec.submodule_search_locations or []) + if spec.origin is not None: + roots.append(str(pathlib.Path(spec.origin).parent)) + for root in roots: + candidate = pathlib.Path(root).resolve() + if (candidate / "include").exists(): + return candidate + + return None + + +@register_dependency("mathdx") +def get_mathdx_include_paths() -> List[str]: + root = get_mathdx_root() + if root is None: + raise RuntimeError( + "Cannot find NVIDIA Math-DX (cuBLASDx) headers. " + "Install the `nvidia-mathdx` package " + "(`pip install nvidia-mathdx`) or set MATHDX_HOME to an " + "extracted Math-DX archive root." + ) + candidates = [root / "include"] + cutlass = root / "external" / "cutlass" / "include" + if cutlass.exists(): + candidates.append(cutlass) + return [str(p) for p in candidates] + + @register_dependency("cutlass") def get_cutlass_include_paths() -> List[str]: include_paths: List[str] = [] diff --git a/python/sglang/srt/eplb/expert_location_dispatch.py b/python/sglang/srt/eplb/expert_location_dispatch.py index 533086e03..261a06c1f 100644 --- a/python/sglang/srt/eplb/expert_location_dispatch.py +++ b/python/sglang/srt/eplb/expert_location_dispatch.py @@ -77,7 +77,9 @@ def transform_select_experts_inputs( def topk_ids_logical_to_physical( - topk_ids: torch.Tensor, info: Optional[ExpertLocationDispatchInfo] + topk_ids: torch.Tensor, + info: Optional[ExpertLocationDispatchInfo], + log2phy_prob: Optional[torch.Tensor] = None, ) -> torch.Tensor: if info is None: return topk_ids @@ -86,6 +88,13 @@ def topk_ids_logical_to_physical( return _topk_ids_logical_to_physical_static(topk_ids, info) if info.ep_dispatch_algorithm in ["dynamic", "fake"]: return _topk_ids_logical_to_physical_dynamic(topk_ids, info) + if info.ep_dispatch_algorithm == "lp": + if log2phy_prob is None: + raise RuntimeError( + "ep_dispatch_algorithm='lp' but log2phy_prob is None at dispatch " + f"time (topk_ids.shape={tuple(topk_ids.shape)})." + ) + return _topk_ids_logical_to_physical_probability(topk_ids, info, log2phy_prob) raise NotImplementedError(f"Unknown algorithm {info.ep_dispatch_algorithm}") @@ -115,3 +124,24 @@ def _topk_ids_logical_to_physical_dynamic( topk_ids = topk_ids.view(topk_ids_original_shape) return topk_ids + + +def _topk_ids_logical_to_physical_probability( + topk_ids: torch.Tensor, + info: ExpertLocationDispatchInfo, + log2phy_prob: torch.Tensor, +) -> torch.Tensor: + """Select physical experts via the JIT-compiled CUDA dispatch kernel. + + Raises if ``topk_ids`` isn't on CUDA — the LP path requires the fused + kernel and there is no torch reference fallback at runtime. + """ + if not topk_ids.is_cuda: + raise RuntimeError( + "LP dispatch requires CUDA tensors; got topk_ids on " f"{topk_ids.device}." + ) + from sglang.jit_kernel.lplb import cuda_solver + + return cuda_solver.dispatch_probability( + topk_ids, log2phy_prob, info.partial_logical_to_all_physical_map + ) diff --git a/python/sglang/srt/eplb/lplb_solver.py b/python/sglang/srt/eplb/lplb_solver.py new file mode 100644 index 000000000..0186a370b --- /dev/null +++ b/python/sglang/srt/eplb/lplb_solver.py @@ -0,0 +1,280 @@ +""" +LPLBSolver — Linear-Programming Load Balancer for Expert Parallelism. + +Encapsulates LP matrix construction (offline, at init/rebalance) and +per-batch solving (online, per MoE layer forward pass). + +Design for DP-attention: + Each EP rank counts its local tokens, then all ranks participate in an + all-reduce to obtain identical global counts. Every rank then solves + the same LP independently, producing the same log2phy_prob — no + broadcast is needed. Empty-token ranks contribute zeros in the + all-reduce so the collective never deadlocks. + +Usage: + solver = LPLBSolver(phy2log, log2phy, num_gpus, ep_group) + log2phy_prob = solver.solve(topk_ids) # per batch +""" + +from __future__ import annotations + +import logging +from typing import Optional + +import torch + +logger = logging.getLogger(__name__) + +# Global per-layer LPLB solvers +_global_lplb_solvers: dict[int, LPLBSolver] = {} + + +# LP dispatch requires every EP rank to call solver.solve() on every forward +# pass (including empty-topk ranks under DP-attention) — the all-reduce inside +# would otherwise hang. Only the DeepSeek-v2 family and its subclasses route +# empty-rank paths through solver.solve(); other MoE families would deadlock. +_LPLB_SUPPORTED_MODEL_ARCHS: frozenset[str] = frozenset( + { + "DeepseekV2ForCausalLM", + "DeepseekV3ForCausalLM", + "DeepseekV32ForCausalLM", + "MistralLarge3ForCausalLM", + "MistralLarge3ForCausalLMEagle", + "Glm4MoeLiteForCausalLM", + "GlmMoeDsaForCausalLM", + } +) + + +def assert_lplb_supported_model(architecture: str) -> None: + if architecture not in _LPLB_SUPPORTED_MODEL_ARCHS: + supported = ", ".join(sorted(_LPLB_SUPPORTED_MODEL_ARCHS)) + raise NotImplementedError( + f"{architecture} does not support --ep-dispatch-algorithm lp. " + f"Validated targets: {supported}. Other MoE families have " + "empty-token early returns that don't participate in the EP " + "all-reduce inside LPLBSolver.solve(), which would deadlock " + "under DP-attention." + ) + + +def get_global_lplb_solver(layer_id: int) -> Optional[LPLBSolver]: + return _global_lplb_solvers.get(layer_id) + + +def set_global_lplb_solver(layer_id: int, solver: LPLBSolver): + _global_lplb_solvers[layer_id] = solver + + +def clear_global_lplb_solvers(): + _global_lplb_solvers.clear() + + +class LPLBSolver: + """ + Per-layer LPLB solver. + + At init: pre-computes LP constraint matrices from expert-to-GPU mapping. + At solve: takes topk_ids, counts tokens, all-reduces, runs LP, + returns log2phy_prob for probability-based token dispatch. + """ + + def __init__( + self, + phy2log: torch.Tensor, + log2phy: torch.Tensor, + num_gpus: int, + ep_group=None, + logical_to_all_physical_map_num_valid=None, + ): + """ + Args: + phy2log: (num_physical_experts,) physical-to-logical expert mapping. + log2phy: (num_logical_experts, max_copies) logical-to-physical mapping (-1 padded). + num_gpus: Number of GPUs in the EP group. + ep_group: GroupCoordinator for EP communication (all-reduce). + logical_to_all_physical_map_num_valid: (num_logical_experts,) number of valid physical copies. + """ + device = phy2log.device + self.num_gpus = num_gpus + self.ep_group = ep_group + self._has_redundancy = False + if logical_to_all_physical_map_num_valid is not None: + self._has_redundancy = bool( + (logical_to_all_physical_map_num_valid > 1).any() + ) + + self.num_logical = log2phy.shape[0] + self.max_copies = log2phy.shape[1] + self.num_phy = phy2log.shape[0] + # B1/B2 GPU-assignment matrices below assume each rank owns a + # contiguous block of num_phy // num_gpus physical experts. + if self.num_phy % num_gpus != 0: + raise ValueError( + f"LPLBSolver requires num_phy ({self.num_phy}) to be divisible " + f"by num_gpus ({num_gpus}); per-rank-contiguous ownership is " + "currently the only supported allocation." + ) + num_phy_per_gpu = self.num_phy // num_gpus + + # Count copies per logical expert + logcnt = torch.bincount(phy2log, minlength=self.num_logical) + + # Separate single-copy vs replicated experts. + # Stored as int64 so they can be used directly as index tensors in + # _solve without per-call .long() casts (Tier 1 optimization). + self.log_single = torch.nonzero(logcnt == 1).flatten().to(torch.int64) + self.phy_single = log2phy[self.log_single, 0].to(torch.int64) + self.log_replicated = torch.nonzero(logcnt > 1).flatten().to(torch.int64) + self.phy_replicated = ( + torch.nonzero(logcnt[phy2log] > 1).flatten().to(torch.int64) + ) + + self.num_single = len(self.log_single) + self.num_red_log = len(self.log_replicated) + self.num_red_phy = len(self.phy_replicated) + + # Build GPU assignment matrices + B_full = torch.zeros( + (num_gpus, self.num_phy), dtype=torch.float32, device=device + ) + for i in range(num_gpus): + B_full[i, i * num_phy_per_gpu : (i + 1) * num_phy_per_gpu] = 1 + self.B1 = B_full[:, self.phy_single].contiguous() + B2 = B_full[:, self.phy_replicated] + + # Build C matrix (copy-to-logical mapping) + C = torch.zeros( + (self.num_red_log, self.num_red_phy), dtype=torch.float32, device=device + ) + phy2log_rep = phy2log[self.phy_replicated] + for i in range(self.num_red_log): + C[i, phy2log_rep == self.log_replicated[i]] = 1.0 + + # Build A_base = [[C, 0, 0], [B2, I, -1]] (without Big-M column) + zeros_top_g = torch.zeros( + (self.num_red_log, num_gpus), dtype=torch.float32, device=device + ) + zeros_top_1 = torch.zeros( + (self.num_red_log, 1), dtype=torch.float32, device=device + ) + I_g = torch.eye(num_gpus, dtype=torch.float32, device=device) + neg_ones = torch.full((num_gpus, 1), -1.0, dtype=torch.float32, device=device) + + A_top = torch.hstack([C, zeros_top_g, zeros_top_1]) + A_bottom = torch.hstack([B2, I_g, neg_ones]) + self.A_base = torch.vstack([A_top, A_bottom]).contiguous() + + # Objective: minimize M (second-to-last var), penalize Big-M auxiliary + nv = self.A_base.shape[1] + 1 # +1 for Big-M column + self.c_vec = torch.zeros(nv, dtype=torch.float32, device=device) + self.c_vec[-2] = 1.0 + self.c_vec[-1] = 1000.0 + + # Store log2phy as int64 so it can be used directly as index tensor + # without per-call .long() casts (Tier 1 optimization). + self.log2phy = log2phy.to(torch.int64).contiguous() + + # Pre-JIT-compile the fused IPM kernel for this (NC, NV) shape so the + # 20-40s compile cost happens once at startup rather than on the first + # real request. No-op when the fused backend is unavailable. + nc = self.A_base.shape[0] + nv = self.A_base.shape[1] + 1 # +1 for Big-M column added in solve() + from sglang.jit_kernel.lplb.torch_solver import warmup as _ipm_warmup + + _ipm_warmup(nc, nv, num_iters=5, device=device) + + # Pre-compute A_base row sum (used in every prep call). + self._A_base_row_sum = self.A_base.sum(dim=1).contiguous() # (NC,) + + # Pre-allocate the buffers the JIT CUDA prep / IPM / post kernels write + # into. All writes are contiguous full-tensor stores (no strided + # ``out=`` semantics), so the reuse is safe under high concurrency. + # Constructed lazily on the first solve() call (we don't know the + # device-side log2phy_prob shape until then) — see _solve. + self._A_full = torch.empty(nc, nv, dtype=torch.float32, device=device) + self._A_full[:, : nv - 1].copy_(self.A_base) + self._b = torch.empty(nc, dtype=torch.float32, device=device) + self._t1 = torch.empty(self.num_single, dtype=torch.float32, device=device) + self._x = torch.empty(nv, dtype=torch.float32, device=device) + self._log2phy_prob = torch.empty( + log2phy.shape, dtype=torch.float32, device=device + ) + + def solve(self, topk_ids: torch.Tensor) -> torch.Tensor: + """ + Full LPLB pipeline: count -> all-reduce -> LP solve -> return log2phy_prob. + + All EP ranks must call this method every MoE layer forward pass, + including empty-token ranks (which pass an empty topk_ids tensor). + This ensures the all-reduce collective does not deadlock under + DP-attention where different ranks may have different token counts. + + Args: + topk_ids: (num_tokens, topk) int32 tensor of logical expert IDs. + Can be empty (shape (0, topk)) for idle ranks. + + Returns: + log2phy_prob: (num_logical, max_copies) float32 probability tensor. + """ + device = topk_ids.device + + # Step 1: Count local tokens per logical expert. + # topk_ids comes from the router and is by construction in + # [0, num_logical), so we can scatter_add directly without filtering. + # Boolean masking + numel() (the previous defensive form) forced a + # GPU->host sync on every forward pass via aten::nonzero and a + # tensor-shape read; scatter_add on the flattened tensor is async + # and a no-op when topk_ids is empty (DP-attention idle rank case). + local_counts = torch.zeros(self.num_logical, dtype=torch.int32, device=device) + flat_ids = topk_ids.flatten() + local_counts.scatter_add_( + 0, + flat_ids.long(), + torch.ones_like(flat_ids, dtype=torch.int32), + ) + + # Step 2: All-reduce to get global counts across all EP ranks. + # All EP ranks must participate — empty-token ranks contribute zeros. + # After all-reduce, every rank has identical global_counts and solves + # the same LP independently, so no broadcast is needed. + # GroupCoordinator.all_reduce may be in-place (pynccl) or out-of-place + # (ca_comm / pymscclpp / ...) depending on tensor size; small tensors + # like ours (~num_logical * 4 B) typically take the out-of-place path, + # so we must capture the return value. + global_counts = local_counts.float() + if self.ep_group is not None: + global_counts = self.ep_group.all_reduce(global_counts) + + # Step 3: Run LP solver + return self._solve(global_counts) + + def _solve(self, global_counts: torch.Tensor) -> torch.Tensor: + """Three CUDA kernel launches replace ~14 torch ops. + + Pipeline (all writes go into pre-allocated buffers from __init__): + prep_lp_inputs → solve_ipm → extract_log2phy_prob + Raises if the JIT CUDA backend is unavailable. + """ + from sglang.jit_kernel.lplb import cuda_solver + + cuda_solver.prep_lp_inputs( + self._A_full, + self._b, + self._t1, + global_counts, + self.log_single, + self.log_replicated, + self.B1, + self._A_base_row_sum, + ) + cuda_solver.solve_ipm(self._A_full, self._b, self.c_vec, result=self._x) + cuda_solver.extract_log2phy_prob( + self._log2phy_prob, + self._x, + self._t1, + self.phy_single, + self.phy_replicated, + self.log2phy, + ) + return self._log2phy_prob diff --git a/python/sglang/srt/layers/moe/hash_topk.py b/python/sglang/srt/layers/moe/hash_topk.py index 1ca160cd2..1f5b93ac1 100644 --- a/python/sglang/srt/layers/moe/hash_topk.py +++ b/python/sglang/srt/layers/moe/hash_topk.py @@ -34,9 +34,10 @@ class HashTopK(nn.Module): scoring_func="sqrtsoftplus", routed_scaling_factor=1.5, apply_routed_scaling_factor_on_output=False, + layer_id: Optional[int] = None, ): super().__init__() - self.layer_id = None + self.layer_id = layer_id from sglang.srt.server_args import get_global_server_args self.enable_deepep_waterfill = ( @@ -80,8 +81,18 @@ class HashTopK(nn.Module): with torch.no_grad(): self.tid2eid.copy_(tid2eid.to(self.tid2eid.dtype)) - def empty_topk_output(self, device: torch.device): + def empty_topk_output( + self, device: torch.device, *, layer_id: Optional[int] = None + ): topk = self.topk - self.num_fused_shared_experts + if layer_id is not None: + from sglang.srt.eplb.lplb_solver import get_global_lplb_solver + + lplb_solver = get_global_lplb_solver(layer_id) + if lplb_solver is not None: + lplb_solver.solve( + torch.empty((0, topk), dtype=torch.int32, device=device) + ) topk_weights = torch.empty((0, topk), dtype=torch.float32, device=device) topk_ids = torch.full((0, topk), -1, dtype=torch.int32, device=device) router_logits = torch.empty((0, topk), dtype=torch.float32, device=device) @@ -175,7 +186,23 @@ class HashTopK(nn.Module): if is_hip(): topk_weights = topk_weights.to(torch.float32) - topk_ids = topk_ids_logical_to_physical(topk_ids, expert_location_dispatch_info) + log2phy_prob = None + if ( + expert_location_dispatch_info is not None + and getattr(expert_location_dispatch_info, "ep_dispatch_algorithm", None) + == "lp" + ): + if self.layer_id is None: + raise RuntimeError("HashTopK LP dispatch requires layer_id.") + from sglang.srt.eplb.lplb_solver import get_global_lplb_solver + + lplb_solver = get_global_lplb_solver(self.layer_id) + if lplb_solver is not None: + log2phy_prob = lplb_solver.solve(topk_ids) + + topk_ids = topk_ids_logical_to_physical( + topk_ids, expert_location_dispatch_info, log2phy_prob + ) if is_hip(): _zero_topk_weights_padded_region(topk_weights, num_token_non_padded) else: diff --git a/python/sglang/srt/layers/moe/topk.py b/python/sglang/srt/layers/moe/topk.py index 3e9bb990a..6e80ee794 100644 --- a/python/sglang/srt/layers/moe/topk.py +++ b/python/sglang/srt/layers/moe/topk.py @@ -552,7 +552,30 @@ class TopK(MultiPlatformOp): layer_id=self.layer_id, ) - def empty_topk_output(self, device: torch.device) -> TopKOutput: + def empty_topk_output( + self, device: torch.device, *, layer_id: Optional[int] = None + ) -> TopKOutput: + """Return an empty topk output for a rank with zero tokens this forward. + + When ``layer_id`` is provided and the active dispatch algorithm is LP, + also calls ``LPLBSolver.solve(empty)`` so that this rank participates + in the EP all-reduce. Without this, an empty rank would skip the + collective and deadlock under DP-attention. + """ + if layer_id is not None: + # Skip the full ExpertLocationDispatchInfo allocation — we only + # need the per-layer solver to participate in the EP all-reduce. + from sglang.srt.eplb.lplb_solver import get_global_lplb_solver + + lplb_solver = get_global_lplb_solver(layer_id) + if lplb_solver is not None: + lplb_solver.solve( + torch.empty( + (0, self.topk_config.top_k), + dtype=torch.int32, + device=device, + ) + ) topk = self.topk_config.top_k - self.topk_config.num_fused_shared_experts with use_symmetric_memory( get_tp_group(), disabled=not is_allocation_symmetric() @@ -1508,11 +1531,29 @@ def _post_process_topk_ids( ) recorder_topk_ids = None if _is_cuda: - # When shared experts are fused (appended as extra columns in topk_ids), - # EPLB dispatch must only remap the routed expert columns. - # The shared expert column (value = n_routed_experts) would be out-of-bounds - # for the logical-to-physical dispatch table. - if num_fused_shared_experts > 0 and is_deepep_class_backend(): + # LP path: solve LP outside torch.compile (the solver contains an + # EP all-reduce that can't run inside compiled regions). + log2phy_prob = None + if ( + expert_location_dispatch_info is not None + and getattr(expert_location_dispatch_info, "ep_dispatch_algorithm", None) + == "lp" + ): + from sglang.srt.eplb.lplb_solver import get_global_lplb_solver + + lplb_solver = get_global_lplb_solver(layer_id) + if lplb_solver is not None: + log2phy_prob = lplb_solver.solve(topk_ids) + + if log2phy_prob is not None: + topk_ids = topk_ids_logical_to_physical( + topk_ids, expert_location_dispatch_info, log2phy_prob + ) + _mask_topk_ids_padded_region(topk_ids, num_token_non_padded) + elif num_fused_shared_experts > 0 and is_deepep_class_backend(): + # Shared experts appended as extra columns in topk_ids: their value + # would be out-of-bounds for the logical-to-physical dispatch table, + # so split, dispatch the routed cols, recombine. shared_cols = topk_ids[:, -num_fused_shared_experts:] routed_cols = topk_ids[:, :-num_fused_shared_experts] routed_cols = _biased_grouped_topk_postprocess( diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index e9ce8579e..5bf823736 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -106,6 +106,12 @@ from sglang.srt.eplb.expert_location import ( set_global_expert_location_metadata, ) from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater +from sglang.srt.eplb.lplb_solver import ( + LPLBSolver, + assert_lplb_supported_model, + clear_global_lplb_solvers, + set_global_lplb_solver, +) from sglang.srt.hardware_backend.npu.graph_runner.npu_graph_runner import NPUGraphRunner from sglang.srt.kv_canary.api import install_canary from sglang.srt.kv_canary.runner.canary_manager import context_tuple @@ -691,6 +697,9 @@ class ModelRunner(ModelRunnerKVCacheMixin): ) ) + if self.server_args.ep_dispatch_algorithm == "lp" and not self.is_draft_worker: + self._init_lplb_solvers() + # Expert parallelism self.eplb_manager = ( EPLBManager(self) @@ -1612,6 +1621,35 @@ class ModelRunner(ModelRunnerKVCacheMixin): logger, f"Prepared {num_prepared} DeepEP waterfill TopK modules." ) + def _init_lplb_solvers(self): + """Initialize per-layer LPLB solvers from current expert location metadata.""" + from sglang.srt.distributed import get_moe_ep_group + + # Gate: refuse LP for non-DeepSeek MoE families whose empty-token paths + # don't participate in the EP all-reduce (would deadlock under DP- + # attention). Failure here happens before any forward pass. + architectures = getattr(self.model_config.hf_config, "architectures", None) + if architectures: + assert_lplb_supported_model(architectures[0]) + + metadata = get_global_expert_location_metadata() + if metadata is None: + return + clear_global_lplb_solvers() + ep_group = get_moe_ep_group() + for lid in range(metadata.num_layers): + solver = LPLBSolver( + phy2log=metadata.physical_to_logical_map[lid], + log2phy=metadata.logical_to_all_physical_map[lid], + num_gpus=metadata.ep_size, + ep_group=ep_group, + logical_to_all_physical_map_num_valid=( + metadata.logical_to_all_physical_map_num_valid[lid] + ), + ) + set_global_lplb_solver(lid, solver) + logger.info(f"Initialized LPLB solvers for {metadata.num_layers} layers") + def update_expert_location( self, new_expert_location_metadata: ExpertLocationMetadata, @@ -1654,6 +1692,10 @@ class ModelRunner(ModelRunnerKVCacheMixin): weight_name_filter=weight_name_filter, ) + # Re-init LPLB solvers after expert location update + if self.server_args.ep_dispatch_algorithm == "lp": + self._init_lplb_solvers() + def maybe_recover_ep_ranks(self): # TODO(perf): `active_ranks.all()` on a CUDA tensor triggers host-device # synchronization, and this function is on the forward-path. diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index d11428646..c6208ff1d 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -631,6 +631,7 @@ class DeepseekV2MoE(nn.Module): scoring_func=config.scoring_func, routed_scaling_factor=self.routed_scaling_factor, apply_routed_scaling_factor_on_output=self.experts.should_fuse_routed_scaling_factor_in_topk, + layer_id=self.layer_id, ) else: # Default: grouped noaux_tc top-k. Covers V3/V3.2/GLM-5/Glm4MoeLite. @@ -996,7 +997,9 @@ class DeepseekV2MoE(nn.Module): ) else: shared_output = None - topk_output = self.topk.empty_topk_output(hidden_states.device) + topk_output = self.topk.empty_topk_output( + hidden_states.device, layer_id=self.layer_id + ) if self._fuse_shared_experts_inside_sbo: shared_output = None @@ -1168,7 +1171,19 @@ class DeepseekV2MoE(nn.Module): **topk_kwargs, ) else: - topk_output = self.topk.empty_topk_output(hidden_states.device) + topk_output = self.topk.empty_topk_output( + hidden_states.device, layer_id=self.layer_id + ) + if is_deepep_class_backend() and self.num_fused_shared_experts > 0: + n = self.num_fused_shared_experts + topk_output = topk_output._replace( + topk_ids=topk_output.topk_ids.new_empty( + (0, topk_output.topk_ids.shape[-1] + n) + ), + topk_weights=topk_output.topk_weights.new_empty( + (0, topk_output.topk_weights.shape[-1] + n) + ), + ) if sbo_overlap_dispatch_flag: shared_output = None @@ -1391,7 +1406,9 @@ class DeepseekV2MoE(nn.Module): ), ) else: - state.topk_output = self.topk.empty_topk_output(hidden_states.device) + state.topk_output = self.topk.empty_topk_output( + hidden_states.device, layer_id=self.layer_id + ) def op_dispatch_a(self, state): if self.ep_size > 1: diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 28cc109de..7f29f3f5e 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -668,7 +668,7 @@ class ServerArgs: "auto" ) ep_num_redundant_experts: int = 0 - ep_dispatch_algorithm: Optional[Literal["static", "dynamic", "fake"]] = None + ep_dispatch_algorithm: Optional[Literal["static", "dynamic", "fake", "lp"]] = None init_expert_location: str = "trivial" enable_eplb: bool = False eplb_algorithm: str = "auto" diff --git a/test/registered/eplb/test_lplb_distributed.py b/test/registered/eplb/test_lplb_distributed.py new file mode 100644 index 000000000..f5996c57e --- /dev/null +++ b/test/registered/eplb/test_lplb_distributed.py @@ -0,0 +1,446 @@ +"""Real-CUDA multi-rank distributed test for LPLBSolver. + +Supplements the CPU-only `test_lplb.py` checks: this test spawns 2 +CUDA-bound ranks on a single node via `torch.multiprocessing.spawn` and +exercises the actual code paths the production LP server uses, including +the EP all-reduce and the DP-attention "empty rank" case (rank with zero +tokens for a given forward). + +Critical invariants this test catches that the CPU tests cannot: + 1. The all-reduce inside `LPLBSolver.solve()` actually completes when + one rank has empty `topk_ids` (i.e., the empty-rank deadlock fix is + in effect — without it, the empty rank would skip `solver.solve()` + and hang the collective). + 2. Both ranks receive the **same** counts post-all-reduce and produce + identical outputs. + 3. The output equals the torch-IPM oracle computed from the ORACLE + summed counts (this is the discriminator that tells "all-reduce + happened" apart from "rank silently used local-only counts and + coincidentally produced a finite tensor"). + 4. The fused-only contract holds end-to-end: a fused-kernel failure + must surface as an exception (no silent torch fallback exists). + 5. Post-rebalance reinit produces solvers tied to the new metadata — + outputs change when `phy2log` changes. + +Pattern follows `test/registered/layers/mamba/test_mamba2_mixer.py` — +`torch.multiprocessing.spawn` is required because `torch.distributed` +plays poorly with arbitrary subprocess launchers. +""" + +import pytest +import torch + +from sglang.srt.distributed.device_communicators.custom_all_reduce_utils import ( + update_environment_variables, +) +from sglang.srt.distributed.parallel_state import ( + init_distributed_environment, + initialize_model_parallel, +) +from sglang.test.ci.ci_register import register_cuda_ci + +register_cuda_ci(est_time=120, stage="base-b", runner_config="2-gpu-large") + +NUM_GPUS = 2 +TOPK = 2 + +# Smallest replicated-metadata config that still exercises a non-trivial LP: +# num_gpus = 2 +# num_logical = 4 (3 single-copy + 1 replicated) +# num_phy = 6 (4 + 2 redundant copies of expert 0) +# → NC = 1 + 2 = 3, NV = 2 + 2 + 2 = 6 (well within H20 shmem budget) +NUM_LOGICAL = 4 +NUM_PHY = 6 + + +def _make_metadata(rebalanced: bool = False): + """Build (phy2log, log2phy, num_valid) for the test config. + + When `rebalanced=True`, swap the redundant logical from expert 0 to + expert 1 — same shapes, different mapping, so a stale solver would + produce wrong outputs. + """ + replicated_logical = 1 if rebalanced else 0 + phy2log = torch.tensor( + [0, 1, 2, 3, replicated_logical, replicated_logical], dtype=torch.int64 + ) + log2phy = torch.full((NUM_LOGICAL, 3), -1, dtype=torch.int64) + for i in range(NUM_LOGICAL): + log2phy[i, 0] = i + log2phy[replicated_logical, 1] = 4 + log2phy[replicated_logical, 2] = 5 + num_valid = torch.ones(NUM_LOGICAL, dtype=torch.int64) + num_valid[replicated_logical] = 3 + return phy2log, log2phy, num_valid + + +@pytest.mark.skipif( + not torch.cuda.is_available(), + reason="This test requires CUDA", +) +def test_dispatch_probability_matches_torch_reference(): + """The fused CUDA `dispatch_probability` and the pure-torch reference + must produce identical outputs for the same ``random_vals``. Single-rank, + runs on any CUDA GPU.""" + from sglang.jit_kernel.lplb.cuda_solver import ( + dispatch_probability, + dispatch_probability_torch_reference, + ) + + torch.manual_seed(0) + device = torch.device("cuda:0") + num_logical, max_copies, num_tokens, topk = 8, 3, 64, 4 + + # Random LP-style probabilities (some rows zeroed to exercise the fallback). + log2phy_prob = torch.rand(num_logical, max_copies, device=device) + log2phy_prob[0] = 0 # force fallback row + log2phy_map = torch.full( + (num_logical, max_copies), -1, dtype=torch.int32, device=device + ) + for i in range(num_logical): + log2phy_map[i, 0] = i # primary replica + log2phy_map[i, 1] = num_logical + (i % 4) # one redundant + # leave last slot at -1 to exercise the masked-replica path + + topk_ids = torch.randint( + 0, num_logical, (num_tokens, topk), dtype=torch.int32, device=device + ) + random_vals = torch.rand(num_tokens * topk, device=device, dtype=torch.float32) + + cuda_out = dispatch_probability(topk_ids, log2phy_prob, log2phy_map, random_vals) + torch_out = dispatch_probability_torch_reference( + topk_ids, log2phy_prob, log2phy_map, random_vals + ) + + assert torch.equal(cuda_out, torch_out), ( + f"dispatch_probability disagrees with torch reference: " + f"{(cuda_out != torch_out).sum().item()}/{cuda_out.numel()} mismatches" + ) + + +@pytest.mark.skipif( + not torch.cuda.is_available(), + reason="This test requires CUDA", +) +def test_solve_ipm_matches_torch_reference(): + """The fused CUDA IPM kernel and the pure-torch reference should agree to + a small tolerance on a *real* LPLB LP. They are NOT bit-equivalent — the + kernel factors the KKT system with a hand-written block Cholesky while the + reference uses torch.linalg.solve (LU) — so we compare with allclose and + print the max abs difference (the numerical-difference number the review + asked about). + + The LP is built the way LPLBSolver does (Big-M column + normalized + counts), because an arbitrary random Ax=b never satisfies the solver's + convergence test (the Big-M slack must reach ~0) and both backends would + just return the 0.5 non-convergence sentinel — agreeing trivially without + exercising the solve. Single-rank, any CUDA GPU.""" + from sglang.jit_kernel.lplb.cuda_solver import solve_ipm as cuda_solve_ipm + from sglang.jit_kernel.lplb.torch_solver import solve_ipm_torch_reference + from sglang.srt.eplb.lplb_solver import LPLBSolver + + torch.manual_seed(0) + device = torch.device("cuda:0") + phy2log, log2phy, num_valid = _make_metadata() + solver = LPLBSolver( + phy2log=phy2log.to(device), + log2phy=log2phy.to(device), + num_gpus=NUM_GPUS, + ep_group=None, + logical_to_all_physical_map_num_valid=num_valid.to(device), + ) + + # Realistic counts: the replicated expert (logical 0) carries heavy load + # to distribute across its 3 physical copies; the rest are moderate. This + # gives the barrier method a well-posed instance that converges in 5 iters. + global_counts = torch.tensor( + [120.0, 30.0, 25.0, 20.0], dtype=torch.float32, device=device + ) + + # Reconstruct the LP exactly as LPLBSolver._solve builds it. + counts_norm = global_counts / global_counts.sum().clamp(min=1.0) + t1 = counts_norm[solver.log_single] + b1 = counts_norm[solver.log_replicated] + b2 = -(solver.B1 @ t1).flatten() + b = torch.cat([b1, b2]) + big_M_col = b - solver._A_base_row_sum + A_full = torch.hstack([solver.A_base, big_M_col.unsqueeze(1)]) + + cuda_x = cuda_solve_ipm(A_full, b, solver.c_vec) + torch_x = solve_ipm_torch_reference(A_full, b, solver.c_vec) + + converged = not torch.allclose(cuda_x, torch.full_like(cuda_x, 0.5)) + max_diff = (cuda_x - torch_x).abs().max().item() + print( + f"\n[ipm-compare] converged={converged} max|cuda-torch|={max_diff:.3e} " + f"cuda={[round(v,4) for v in cuda_x.tolist()]} " + f"torch={[round(v,4) for v in torch_x.tolist()]}" + ) + assert converged, ( + "IPM returned the 0.5 non-convergence sentinel — the comparison would " + "be trivial. Adjust the LP instance so it converges." + ) + assert torch.allclose( + cuda_x, torch_x, atol=1e-2, rtol=1e-2 + ), f"fused IPM diverges from torch reference: max abs diff {max_diff:.3e}" + + +@pytest.mark.skipif( + not torch.cuda.is_available() or torch.cuda.device_count() < NUM_GPUS, + reason=f"This test requires at least {NUM_GPUS} CUDA devices", +) +def test_lplb_distributed_two_rank(): + """Driver: spawn `NUM_GPUS` worker processes that each run the + distributed-LPLB invariant suite.""" + torch.multiprocessing.spawn( + _worker_main, + args=(NUM_GPUS,), + nprocs=NUM_GPUS, + ) + + +def _worker_main(local_rank: int, world_size: int): + """Per-rank entry point under torch.multiprocessing.spawn.""" + # Inject minimal ServerArgs before any LPLB module reads the global state. + # Silent fallbacks no longer exist — the fused CUDA path is the only LP + # path, so this test relies on hard failures, not gating flags. + from sglang.srt.server_args import ( + ServerArgs, + set_global_server_args_for_scheduler, + ) + + set_global_server_args_for_scheduler( + ServerArgs( + model_path="dummy", + ) + ) + + device = torch.device(f"cuda:{local_rank}") + torch.cuda.set_device(device) + torch.set_default_device(device) + + update_environment_variables( + { + "RANK": str(local_rank), + "LOCAL_RANK": str(local_rank), + "WORLD_SIZE": str(world_size), + "MASTER_ADDR": "localhost", + "MASTER_PORT": "12347", # Distinct from other tests' ports. + } + ) + init_distributed_environment( + world_size=world_size, rank=local_rank, local_rank=local_rank + ) + initialize_model_parallel( + tensor_model_parallel_size=world_size, + expert_model_parallel_size=world_size, + ) + + from sglang.srt.eplb.lplb_solver import clear_global_lplb_solvers + + try: + # Clear the global solver registry between subtests so a stale solver + # can't shadow a freshly-built one. + clear_global_lplb_solvers() + _check_solver_with_empty_rank(local_rank, world_size, device) + clear_global_lplb_solvers() + _check_all_ranks_empty(local_rank, world_size, device) + clear_global_lplb_solvers() + _check_solver_determinism(local_rank, world_size, device) + clear_global_lplb_solvers() + _check_post_rebalance_reinit(local_rank, world_size, device) + finally: + # Use parallel_state's destroy helpers rather than raw + # `dist.destroy_process_group()` so we don't leave _WORLD/_MOE_EP/_TP + # globals stale for later tests. + from sglang.srt.distributed.parallel_state import ( + destroy_distributed_environment, + destroy_model_parallel, + ) + + clear_global_lplb_solvers() + destroy_model_parallel() + destroy_distributed_environment() + + +def _build_solver(): + """Construct an LPLBSolver bound to the live moe_ep_group.""" + from sglang.srt.distributed.parallel_state import get_moe_ep_group + from sglang.srt.eplb.lplb_solver import LPLBSolver + + phy2log, log2phy, num_valid = _make_metadata(rebalanced=False) + return ( + LPLBSolver( + phy2log=phy2log.cuda(), + log2phy=log2phy.cuda(), + num_gpus=NUM_GPUS, + ep_group=get_moe_ep_group(), + logical_to_all_physical_map_num_valid=num_valid.cuda(), + ), + phy2log, + log2phy, + num_valid, + ) + + +def _expected_output(solver, expected_global_counts: torch.Tensor) -> torch.Tensor: + """Oracle: what `solver.solve(*)` should produce given known global + counts. Uses the solver's own `_solve` so we match implementation + quirks (regularization, on-device clamp normalization, etc.). + """ + return solver._solve(expected_global_counts.float().cuda()) + + +def _check_solver_with_empty_rank(rank: int, world_size: int, device: torch.device): + """Rank 0 has real `topk_ids`; rank 1 has `(0, topk)` empty. + + The all-reduce inside `solve()` must complete on both ranks, and both + must produce output equal to the oracle solver-output for the SUMMED + global counts (rank 0's bincount + rank 1's zero contribution). + """ + solver, _, _, _ = _build_solver() + + # Rank 0 owns these tokens; rank 1 is idle (DP-attention empty case). + rank0_topk = torch.tensor( + [[0, 1], [0, 2], [3, 0], [0, 1]], dtype=torch.int32, device=device + ) + if rank == 0: + topk_ids = rank0_topk + else: + topk_ids = torch.empty((0, TOPK), dtype=torch.int32, device=device) + + # Compute the oracle: the global counts are exactly the bincount of + # rank 0's tokens (rank 1 contributes zero). + flat = rank0_topk.flatten().long() + expected_counts = torch.bincount(flat, minlength=NUM_LOGICAL) + expected = _expected_output(solver, expected_counts) + + # Run the actual collective. + actual = solver.solve(topk_ids) + + # Shape + finiteness sanity. + assert actual.shape == ( + NUM_LOGICAL, + solver.max_copies, + ), f"rank {rank}: bad output shape {actual.shape}" + assert torch.isfinite(actual).all(), f"rank {rank}: non-finite values in output" + assert (actual >= 0).all(), f"rank {rank}: negative probabilities" + + # The strong invariant: output matches the summed-count oracle. + # If the all-reduce silently failed, rank 1 would feed all-zero local + # counts to its solver, the oracle wouldn't match, and this assertion + # would catch it. + assert torch.allclose(actual, expected, atol=1e-4, rtol=1e-3), ( + f"rank {rank}: solve() output disagrees with summed-count oracle " + f"(max abs diff {(actual - expected).abs().max().item():.2e}). " + "This usually indicates the EP all-reduce inside solve() was " + "skipped on the empty rank, or rank 1 used local counts only." + ) + + # 0.5-uniform sentinel detection: torch IPM's failure path returns a + # vector of all-0.5. If we see that, the IPM diverged silently. + is_uniform_half = torch.allclose( + actual, + torch.full_like(actual, 0.5), + atol=1e-6, + ) + assert not is_uniform_half, ( + f"rank {rank}: output is the 0.5-uniform fallback, indicating " + "torch IPM convergence failure" + ) + + +def _check_all_ranks_empty(rank: int, world_size: int, device: torch.device): + """Both ranks have empty topk_ids (idle batch). The collective must + still complete; the output should be the all-zero-counts oracle.""" + solver, _, _, _ = _build_solver() + + empty = torch.empty((0, TOPK), dtype=torch.int32, device=device) + expected_counts = torch.zeros(NUM_LOGICAL, dtype=torch.int64) + expected = _expected_output(solver, expected_counts) + + actual = solver.solve(empty) + + assert torch.isfinite(actual).all(), f"rank {rank}: non-finite for empty-batch" + assert (actual >= 0).all(), f"rank {rank}: negative for empty-batch" + assert torch.allclose( + actual, expected, atol=1e-4, rtol=1e-3 + ), f"rank {rank}: empty-batch output disagrees with all-zero oracle" + + +def _check_solver_determinism(rank: int, world_size: int, device: torch.device): + """Same input, two `solve()` calls in a row → identical output.""" + solver, _, _, _ = _build_solver() + if rank == 0: + topk_ids = torch.tensor([[0, 1], [3, 0]], dtype=torch.int32, device=device) + else: + topk_ids = torch.empty((0, TOPK), dtype=torch.int32, device=device) + + out1 = solver.solve(topk_ids) + out2 = solver.solve(topk_ids) + assert torch.equal(out1, out2), ( + f"rank {rank}: solve() not deterministic across calls " + f"(max diff {(out1 - out2).abs().max().item():.2e})" + ) + + +def _check_post_rebalance_reinit(rank: int, world_size: int, device: torch.device): + """Build a NEW solver with different `phy2log` (the redundant logical + moves from expert 0 to expert 1). Same input → DIFFERENT output, since + the LP problem changed.""" + from sglang.srt.distributed.parallel_state import get_moe_ep_group + from sglang.srt.eplb.lplb_solver import LPLBSolver + + solver_old, _, _, _ = _build_solver() + phy2log_new, log2phy_new, num_valid_new = _make_metadata(rebalanced=True) + solver_new = LPLBSolver( + phy2log=phy2log_new.cuda(), + log2phy=log2phy_new.cuda(), + num_gpus=NUM_GPUS, + ep_group=get_moe_ep_group(), + logical_to_all_physical_map_num_valid=num_valid_new.cuda(), + ) + + # An input that hits the replicated logical in BOTH configs (so the + # LP solution is sensitive to which logical is the replicated one). + rank0_topk = torch.tensor( + [[0, 1], [0, 1], [0, 1], [2, 3]], dtype=torch.int32, device=device + ) + if rank == 0: + topk_ids = rank0_topk + else: + topk_ids = torch.empty((0, TOPK), dtype=torch.int32, device=device) + + out_old = solver_old.solve(topk_ids) + out_new = solver_new.solve(topk_ids) + + # The two outputs should differ — same `topk_ids`, different metadata + # ⇒ different LP solution. If they're identical, one of the solvers + # used stale state. + assert not torch.equal(out_old, out_new), ( + f"rank {rank}: post-rebalance solver produced same output as " + "pre-rebalance solver despite different phy2log mappings — " + "stale solver state likely" + ) + + # Positive oracle: each output should match its OWN solver's + # `_solve` of the SUMMED global counts (rank 0's bincount; rank 1 + # contributes zero). Both ranks compare against the same global oracle + # since solve() all-reduces the counts. + expected_counts = torch.bincount(rank0_topk.flatten().long(), minlength=NUM_LOGICAL) + expected_old = _expected_output(solver_old, expected_counts) + expected_new = _expected_output(solver_new, expected_counts) + assert torch.allclose(out_old, expected_old, atol=1e-4, rtol=1e-3), ( + f"rank {rank}: pre-rebalance output disagrees with own oracle " + f"(diff {(out_old - expected_old).abs().max().item():.2e})" + ) + assert torch.allclose(out_new, expected_new, atol=1e-4, rtol=1e-3), ( + f"rank {rank}: post-rebalance output disagrees with own oracle " + f"(diff {(out_new - expected_new).abs().max().item():.2e})" + ) + + +if __name__ == "__main__": + import sys + + sys.exit(pytest.main([__file__, "-v"]))