[DeepSeek V3] Defer moe finalize and fused it with main stream add (#27720)

Co-authored-by: Qiaolin Yu <liin1211@outlook.com>
This commit is contained in:
Khoa Pham
2026-06-12 19:00:51 -07:00
committed by GitHub
co-authored by Qiaolin Yu
parent a14d1a5656
commit d1a39b0c74
7 changed files with 743 additions and 36 deletions
@@ -0,0 +1,418 @@
// Copyright (c) 2026 LightSeek Foundation
//
// Permission is hereby granted, free of charge, to any person obtaining a copy
// of this software and associated documentation files (the "Software"), to deal
// in the Software without restriction, including without limitation the rights
// to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
// copies of the Software, and to permit persons to whom the Software is
// furnished to do so, subject to the following conditions:
//
// The above copyright notice and this permission notice shall be included in
// all copies or substantial portions of the Software.
//
// THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
// IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
// FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
// AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
// LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
// OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
// SOFTWARE.
/*
* Fused MoE finalize + shared-output add (bf16 output, SM>=90 for PDL).
*
* Forked from flashinfer's ``finalizeKernel`` and ``finalizeKernelVecLoad``
* (trtllm_fused_moe_dev_kernel.cu:639 and :803), stripped of the MoE
* backend's KernelParams / UsePdl templating, and extended with an
* optional shared_output residual add on the epilogue side.
*
* For each token t, computes:
* out[t] = Σ_k expert_weights[t, k] * gemm2_out[permuted_idx(t, k)]
* + shared_output[t] // if non-null
*
* Eliminates the native PyTorch ``routed + shared_output`` add (and the
* separate ``*= routed_scaling_factor`` kernel when applicable) from
* ``DeepseekV3MoE.forward``, and gives the downstream allreduce+rmsnorm
* a clean PDL handoff.
*
* Expert-weight dtype is templated on ``TypeExpW`` so we support both the
* bf16 and fp32 topk-weight paths (DSv3/K2.5 trtllm backends use fp32
* because their ``_routing_logits_dtype = torch.float32``; other backends
* use bf16).
*
* Expert-weight scale convention: in our target backends
* (flashinfer trtllm nvfp4 + unquantized), ``apply_routed_scaling_factor_on_output``
* is True, so the routed scaling factor is already folded into
* ``expert_weights`` at topk time. This kernel does not apply any
* additional scale.
*/
#include <cutlass/array.h>
#include <cutlass/numeric_conversion.h>
#include <cutlass/numeric_types.h>
#include "tvm_ffi_utils.h"
#include <cuda_runtime.h>
namespace sglang {
using BF16 = cutlass::bfloat16_t;
constexpr int FINALIZE_THREADS_PER_BLOCK = 256;
constexpr int MAX_TOPK = 64;
// ---------------------------------------------------------------------------
// General kernel — one CTA per (hidden_chunk, token). Picks up small-to-mid
// workloads where the block count fits in a few waves.
// ---------------------------------------------------------------------------
template <typename TypeExpW>
__global__ void moeFinalizeKernel(
int numTokens,
int hiddenDim,
int hiddenDimPadded,
int topK,
BF16 const* __restrict__ inPtr,
int const* __restrict__ expandedIdxToPermutedIdx,
TypeExpW const* __restrict__ expertWeightsPtr,
BF16 const* __restrict__ sharedBiasPtr,
BF16* __restrict__ outPtr) {
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
cudaGridDependencySynchronize();
#endif
for (int64_t tokenIdx = blockIdx.y; tokenIdx < numTokens; tokenIdx += gridDim.y) {
for (int64_t hiddenIdx = threadIdx.x + blockDim.x * blockIdx.x; hiddenIdx < hiddenDim;
hiddenIdx += blockDim.x * gridDim.x) {
float acc = 0.0f;
for (int k = 0; k < topK; k++) {
int64_t const expandedIdx = tokenIdx * topK + k;
int64_t const permutedIdx = expandedIdxToPermutedIdx[expandedIdx];
if (permutedIdx == -1) {
continue;
}
float const scale = static_cast<float>(expertWeightsPtr[expandedIdx]);
float const val = static_cast<float>(inPtr[permutedIdx * hiddenDimPadded + hiddenIdx]);
acc += scale * val;
}
if (sharedBiasPtr != nullptr) {
acc += static_cast<float>(sharedBiasPtr[tokenIdx * hiddenDim + hiddenIdx]);
}
outPtr[tokenIdx * hiddenDim + hiddenIdx] = static_cast<BF16>(acc);
}
}
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
cudaTriggerProgrammaticLaunchCompletion();
#endif
}
// ---------------------------------------------------------------------------
// Vectorized-load kernel — one CTA per token, 128-bit loads, topK unrolled.
// Better at prefill shapes where the general kernel's block count saturates
// many waves and the indirect gather from gemm2_out dominates.
// ---------------------------------------------------------------------------
__device__ inline float4 vectorizedLoadPtx(float4 const* ptr) {
float4 ret;
asm volatile("ld.global.v4.f32 {%0, %1, %2, %3}, [%4];"
: "=f"(ret.x), "=f"(ret.y), "=f"(ret.z), "=f"(ret.w)
: "l"(ptr));
return ret;
}
template <int TopKUnrollFactor>
struct IdxPackedTraits;
template <>
struct IdxPackedTraits<1> {
using Packed = int;
};
template <>
struct IdxPackedTraits<2> {
using Packed = int2;
};
template <>
struct IdxPackedTraits<4> {
using Packed = int4;
};
template <typename TypeExpW, int TopKUnrollFactor>
__global__ void moeFinalizeKernelVecLoad(
int numTokens,
int hiddenDim,
int hiddenDimPadded,
int topK,
BF16 const* __restrict__ inPtr,
int const* __restrict__ expandedIdxToPermutedIdx,
TypeExpW const* __restrict__ expertWeightsPtr,
BF16 const* __restrict__ sharedBiasPtr,
BF16* __restrict__ outPtr) {
static_assert(
TopKUnrollFactor == 1 || TopKUnrollFactor == 2 || TopKUnrollFactor == 4, "TopKUnrollFactor must be 1, 2, or 4");
using IdxPackedType = typename IdxPackedTraits<TopKUnrollFactor>::Packed;
using IdxArrayType = cutlass::Array<int, TopKUnrollFactor>;
using ScaleArrayType = cutlass::Array<TypeExpW, TopKUnrollFactor>;
// 128 bits per thread → 8 bf16 elements.
constexpr int FINALIZE_ELEM_PER_THREAD = 8;
using InputElem = cutlass::Array<BF16, FINALIZE_ELEM_PER_THREAD>;
using OutputElem = cutlass::Array<BF16, FINALIZE_ELEM_PER_THREAD>;
using ComputeElem = cutlass::Array<float, FINALIZE_ELEM_PER_THREAD>;
int64_t const tokenIdx = blockIdx.x;
int64_t const startOffset = threadIdx.x;
int64_t const stride = FINALIZE_THREADS_PER_BLOCK;
int64_t const numElemsInPaddedCol = hiddenDimPadded / FINALIZE_ELEM_PER_THREAD;
int64_t const numElemsInCol = hiddenDim / FINALIZE_ELEM_PER_THREAD;
// Stage the per-token (topK/unroll) indices + scales into smem.
__shared__ ScaleArrayType scaleArrSmem[MAX_TOPK / TopKUnrollFactor];
__shared__ IdxArrayType permutedIdxArrSmem[MAX_TOPK / TopKUnrollFactor];
for (int kChunkIdx = threadIdx.x; kChunkIdx < topK / TopKUnrollFactor; kChunkIdx += blockDim.x) {
int64_t const expandedIdx = tokenIdx * topK + kChunkIdx * TopKUnrollFactor;
auto const permutedIdxPacked =
reinterpret_cast<IdxPackedType const*>(expandedIdxToPermutedIdx)[expandedIdx / TopKUnrollFactor];
permutedIdxArrSmem[kChunkIdx] = *reinterpret_cast<IdxArrayType const*>(&permutedIdxPacked);
#pragma unroll
for (int ki = 0; ki < TopKUnrollFactor; ++ki) {
scaleArrSmem[kChunkIdx][ki] = expertWeightsPtr[expandedIdx + ki];
}
}
BF16* outputPtr = outPtr + tokenIdx * hiddenDim;
auto* outElemPtr = reinterpret_cast<OutputElem*>(outputPtr);
auto const* inElemPtr = reinterpret_cast<InputElem const*>(inPtr);
auto const* sharedElemPtr =
sharedBiasPtr != nullptr ? reinterpret_cast<InputElem const*>(sharedBiasPtr + tokenIdx * hiddenDim) : nullptr;
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
cudaGridDependencySynchronize();
#endif
__syncthreads();
for (int elemIndex = startOffset; elemIndex < numElemsInCol; elemIndex += stride) {
ComputeElem threadOutput;
threadOutput.fill(0.0f);
for (int kChunkIdx = 0; kChunkIdx < topK / TopKUnrollFactor; kChunkIdx++) {
IdxArrayType permutedIdxArr = permutedIdxArrSmem[kChunkIdx];
InputElem inputElemArr[TopKUnrollFactor];
#pragma unroll
for (int ki = 0; ki < TopKUnrollFactor; ++ki) {
int const permutedIdx = permutedIdxArr[ki];
if (permutedIdx == -1) {
continue;
}
auto const* inputPermutedPtr = inElemPtr + permutedIdx * numElemsInPaddedCol;
float4 input = vectorizedLoadPtx(reinterpret_cast<float4 const*>(&inputPermutedPtr[elemIndex]));
inputElemArr[ki] = *reinterpret_cast<InputElem const*>(&input);
}
ScaleArrayType scaleArr = scaleArrSmem[kChunkIdx];
#pragma unroll
for (int ki = 0; ki < TopKUnrollFactor; ++ki) {
int const permutedIdx = permutedIdxArr[ki];
if (permutedIdx == -1) {
continue;
}
float const scale = static_cast<float>(scaleArr[ki]);
cutlass::NumericArrayConverter<float, BF16, FINALIZE_ELEM_PER_THREAD> toFloat;
ComputeElem expertResult = toFloat(inputElemArr[ki]);
#pragma unroll
for (int e = 0; e < FINALIZE_ELEM_PER_THREAD; ++e) {
threadOutput[e] += scale * expertResult[e];
}
}
}
if (sharedElemPtr != nullptr) {
float4 shared = vectorizedLoadPtx(reinterpret_cast<float4 const*>(&sharedElemPtr[elemIndex]));
InputElem sharedElem = *reinterpret_cast<InputElem const*>(&shared);
cutlass::NumericArrayConverter<float, BF16, FINALIZE_ELEM_PER_THREAD> toFloat;
ComputeElem sharedFloat = toFloat(sharedElem);
#pragma unroll
for (int e = 0; e < FINALIZE_ELEM_PER_THREAD; ++e) {
threadOutput[e] += sharedFloat[e];
}
}
cutlass::NumericArrayConverter<BF16, float, FINALIZE_ELEM_PER_THREAD> toBF16;
outElemPtr[elemIndex] = toBF16(threadOutput);
}
#if defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 900)
cudaTriggerProgrammaticLaunchCompletion();
#endif
}
// ---------------------------------------------------------------------------
// Typed dispatch
// ---------------------------------------------------------------------------
template <typename TypeExpW>
void dispatchFinalize(
int numTokens,
int hiddenDim,
int hiddenDimPadded,
int topK,
BF16 const* inPtr,
int const* expandedIdxPtr,
void const* weightsPtrVoid,
BF16 const* sharedPtr,
BF16* outPtr,
bool useVecLoad,
cudaStream_t stream,
cudaLaunchAttribute const* attrs,
int numAttrs) {
auto const* weightsPtr = static_cast<TypeExpW const*>(weightsPtrVoid);
constexpr int kNumThreads = 256;
if (!useVecLoad) {
int const numBlocksX = (hiddenDim + kNumThreads - 1) / kNumThreads;
int const numBlocksY = std::min(8192, numTokens);
cudaLaunchConfig_t config;
config.gridDim = dim3(numBlocksX, numBlocksY);
config.blockDim = dim3(kNumThreads);
config.dynamicSmemBytes = 0;
config.stream = stream;
config.numAttrs = numAttrs;
config.attrs = const_cast<cudaLaunchAttribute*>(attrs);
cudaLaunchKernelEx(
&config,
moeFinalizeKernel<TypeExpW>,
numTokens,
hiddenDim,
hiddenDimPadded,
topK,
inPtr,
expandedIdxPtr,
weightsPtr,
sharedPtr,
outPtr);
return;
}
auto launch = [&](auto unroll_tag) {
constexpr int UNROLL = decltype(unroll_tag)::value;
cudaLaunchConfig_t config;
config.gridDim = dim3(numTokens);
config.blockDim = dim3(FINALIZE_THREADS_PER_BLOCK);
config.dynamicSmemBytes = 0;
config.stream = stream;
config.numAttrs = numAttrs;
config.attrs = const_cast<cudaLaunchAttribute*>(attrs);
cudaLaunchKernelEx(
&config,
moeFinalizeKernelVecLoad<TypeExpW, UNROLL>,
numTokens,
hiddenDim,
hiddenDimPadded,
topK,
inPtr,
expandedIdxPtr,
weightsPtr,
sharedPtr,
outPtr);
};
// Match flashinfer's LAUNCH_TOPK_EXPW dispatch order.
if (topK % 4 == 0) {
launch(std::integral_constant<int, 4>{});
} else if (topK % 2 == 0) {
launch(std::integral_constant<int, 2>{});
} else {
launch(std::integral_constant<int, 1>{});
}
}
} // namespace sglang
// ---------------------------------------------------------------------------
// Host launcher
// ---------------------------------------------------------------------------
void moe_finalize_fuse_shared(
TensorView out,
TensorView gemm2_out,
TensorView expanded_idx_to_permuted_idx,
TensorView expert_weights,
TensorView shared_output,
int64_t top_k,
bool enable_pdl) {
TVM_FFI_ICHECK_EQ(out.ndim(), 2) << "out must be 2-D [numTokens, hiddenDim]";
TVM_FFI_ICHECK_EQ(gemm2_out.ndim(), 2) << "gemm2_out must be 2-D [totalNumPaddedTokens, hiddenDimPadded]";
TVM_FFI_ICHECK_EQ(expanded_idx_to_permuted_idx.ndim(), 1);
TVM_FFI_ICHECK_EQ(expert_weights.ndim(), 2) << "expert_weights must be 2-D [numTokens, topK]";
int const numTokens = int(out.size(0));
int const hiddenDim = int(out.size(1));
int const hiddenDimPadded = int(gemm2_out.size(1));
TVM_FFI_ICHECK_LE(top_k, sglang::MAX_TOPK);
TVM_FFI_ICHECK_EQ(expanded_idx_to_permuted_idx.size(0), numTokens * top_k);
TVM_FFI_ICHECK_EQ(expert_weights.size(0), numTokens);
TVM_FFI_ICHECK_EQ(expert_weights.size(1), top_k);
bool const hasShared = shared_output.numel() > 0;
if (hasShared) {
TVM_FFI_ICHECK_EQ(shared_output.ndim(), 2);
TVM_FFI_ICHECK_EQ(shared_output.size(0), numTokens);
TVM_FFI_ICHECK_EQ(shared_output.size(1), hiddenDim);
}
auto const* inPtr = static_cast<sglang::BF16 const*>(gemm2_out.data_ptr());
auto const* expandedIdxPtr = static_cast<int const*>(expanded_idx_to_permuted_idx.data_ptr());
auto const* sharedPtr = hasShared ? static_cast<sglang::BF16 const*>(shared_output.data_ptr()) : nullptr;
auto* outPtr = static_cast<sglang::BF16*>(out.data_ptr());
cudaSetDevice(out.device().device_id);
cudaStream_t const stream = get_stream(out.device());
// Dispatch heuristic (matches flashinfer): few waves → general kernel,
// many waves → vectorized. The 1184 threshold comes from 148 SMs × 8
// blocks/SM on Blackwell.
constexpr int kNumThreads = 256;
int const numBlocksX = (hiddenDim + kNumThreads - 1) / kNumThreads;
int const numBlocksY = std::min(8192, numTokens);
bool const useVecLoad = (numBlocksX * numBlocksY) >= 1184 && (hiddenDim % 8 == 0) && (hiddenDimPadded % 8 == 0);
cudaLaunchAttribute attrs[1];
attrs[0].id = cudaLaunchAttributeProgrammaticStreamSerialization;
attrs[0].val.programmaticStreamSerializationAllowed = enable_pdl;
auto ew_dtype = expert_weights.dtype();
if (ew_dtype == DLDataType{kDLFloat, 32, 1}) {
sglang::dispatchFinalize<float>(
numTokens,
hiddenDim,
hiddenDimPadded,
int(top_k),
inPtr,
expandedIdxPtr,
expert_weights.data_ptr(),
sharedPtr,
outPtr,
useVecLoad,
stream,
attrs,
1);
} else if (ew_dtype == DLDataType{kDLBfloat, 16, 1}) {
sglang::dispatchFinalize<sglang::BF16>(
numTokens,
hiddenDim,
hiddenDimPadded,
int(top_k),
inPtr,
expandedIdxPtr,
expert_weights.data_ptr(),
sharedPtr,
outPtr,
useVecLoad,
stream,
attrs,
1);
} else {
TVM_FFI_ICHECK(false) << "expert_weights dtype must be float32 or bfloat16";
}
cudaError_t const err = cudaGetLastError();
TVM_FFI_ICHECK(err == cudaSuccess) << "moe_finalize_fuse_shared launch failed: " << cudaGetErrorString(err);
}
TVM_FFI_DLL_EXPORT_TYPED_FUNC(moe_finalize_fuse_shared, moe_finalize_fuse_shared);
@@ -0,0 +1,105 @@
/*
* Copyright (c) 2023 by FlashInfer team.
*
* Licensed under the Apache License, Version 2.0 (the "License");
* you may not use this file except in compliance with the License.
* You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing, software
* distributed under the License is distributed on an "AS IS" BASIS,
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
* See the License for the specific language governing permissions and
* limitations under the License.
*/
#pragma once
#include <tvm/ffi/container/tensor.h>
#include <tvm/ffi/dtype.h>
#include <tvm/ffi/error.h>
#include <tvm/ffi/extra/c_env_api.h>
#include <tvm/ffi/function.h>
#include "dlpack/dlpack.h"
using tvm::ffi::Tensor;
using tvm::ffi::TensorView;
namespace ffi = tvm::ffi;
inline constexpr int64_t encode_dlpack_dtype(DLDataType dtype) {
return (dtype.code << 16) | (dtype.bits << 8) | dtype.lanes;
}
constexpr DLDataType dl_uint8 = DLDataType{kDLUInt, 8, 1};
constexpr DLDataType dl_uint16 = DLDataType{kDLUInt, 16, 1};
constexpr DLDataType dl_uint32 = DLDataType{kDLUInt, 32, 1};
constexpr DLDataType dl_uint64 = DLDataType{kDLUInt, 64, 1};
constexpr DLDataType dl_int8 = DLDataType{kDLInt, 8, 1};
constexpr DLDataType dl_int16 = DLDataType{kDLInt, 16, 1};
constexpr DLDataType dl_int32 = DLDataType{kDLInt, 32, 1};
constexpr DLDataType dl_int64 = DLDataType{kDLInt, 64, 1};
constexpr DLDataType dl_float16 = DLDataType{kDLFloat, 16, 1};
constexpr DLDataType dl_float32 = DLDataType{kDLFloat, 32, 1};
constexpr DLDataType dl_float64 = DLDataType{kDLFloat, 64, 1};
constexpr DLDataType dl_float8_e4m3fn = DLDataType{kDLFloat8_e4m3fn, 8, 1};
constexpr DLDataType dl_float8_e5m2 = DLDataType{kDLFloat8_e5m2, 8, 1};
constexpr DLDataType dl_float4_e2m1fn = DLDataType{kDLFloat4_e2m1fn, 4, 1};
constexpr DLDataType dl_float4_e2m1fn_x2 = DLDataType{kDLFloat4_e2m1fn, 4, 2};
constexpr DLDataType dl_bfloat16 = DLDataType{kDLBfloat, 16, 1};
constexpr DLDataType dl_bool = DLDataType{kDLBool, 8, 1};
constexpr int64_t float16_code = encode_dlpack_dtype(dl_float16);
constexpr int64_t bfloat16_code = encode_dlpack_dtype(dl_bfloat16);
constexpr int64_t float32_code = encode_dlpack_dtype(dl_float32);
constexpr int64_t uint8_code = encode_dlpack_dtype(dl_uint8);
constexpr int64_t int32_code = encode_dlpack_dtype(dl_int32);
constexpr int64_t int64_code = encode_dlpack_dtype(dl_int64);
constexpr int64_t float8_e4m3fn_code = encode_dlpack_dtype(dl_float8_e4m3fn);
constexpr int64_t float8_e5m2_code = encode_dlpack_dtype(dl_float8_e5m2);
constexpr int64_t float4_e2m1fn_code = encode_dlpack_dtype(dl_float4_e2m1fn);
constexpr DLDevice cpu = DLDevice{kDLCPU, 0};
#define CHECK_CUDA(x) TVM_FFI_ICHECK_EQ(x.device().device_type, kDLCUDA) << #x " must be a CUDA tensor";
#define CHECK_CPU(x) TVM_FFI_ICHECK_EQ(x.device().device_type, kDLCPU) << #x " must be a host tensor";
#define CHECK_CONTIGUOUS(x) TVM_FFI_ICHECK(x.IsContiguous()) << #x " must be contiguous";
#define CHECK_LAST_DIM_CONTIGUOUS(x) \
TVM_FFI_ICHECK_EQ(x.stride(-1), 1) \
#x "must be contiguous at last dimension";
#define CHECK_INPUT(x) \
CHECK_CUDA(x); \
CHECK_CONTIGUOUS(x)
#define CHECK_INPUT_TYPE(x, st) TVM_FFI_ICHECK_EQ(x.dtype(), st) << "Inconsistency of Tensor type: " #x;
#define CHECK_INPUT_AND_TYPE(x, st) \
CHECK_CUDA(x); \
CHECK_CONTIGUOUS(x); \
CHECK_INPUT_TYPE(x, st)
#define CHECK_LAST_DIM_CONTIGUOUS_INPUT(x) \
CHECK_CUDA(x); \
CHECK_LAST_DIM_CONTIGUOUS(x)
#define CHECK_DIM(d, x) TVM_FFI_ICHECK_EQ(x.ndim(), d) << #x " must be a " #d "D tensor";
#define CHECK_DEVICE(a, b) \
TVM_FFI_ICHECK_EQ(a.device().device_type, b.device().device_type); \
TVM_FFI_ICHECK_EQ(a.device().device_id, b.device().device_id);
inline cudaStream_t get_current_stream() {
int device;
cudaGetDevice(&device);
return static_cast<cudaStream_t>(TVMFFIEnvGetStream(kDLCUDA, device));
}
inline cudaStream_t get_stream(DLDevice device) {
return static_cast<cudaStream_t>(TVMFFIEnvGetStream(device.device_type, device.device_id));
}
inline int64_t get_element_size(ffi::Tensor x) {
return (x.dtype().bits * x.dtype().lanes) / 8;
}
inline int64_t get_element_size(ffi::TensorView x) {
return (x.dtype().bits * x.dtype().lanes) / 8;
}
inline ffi::Tensor alloc_tensor(tvm::ffi::Shape shape, DLDataType dtype, DLDevice device) {
return ffi::Tensor::FromEnvAlloc(TVMFFIEnvTensorAlloc, shape, dtype, device);
}
@@ -0,0 +1,60 @@
from __future__ import annotations
from typing import Optional
import torch
from sglang.jit_kernel.utils import cache_once, load_jit
@cache_once
def _jit_module():
return load_jit(
"moe_finalize_fuse_shared",
cuda_files=["moe/moe_finalize_fuse_shared.cu"],
extra_dependencies=["cutlass"],
header_only=False,
)
def moe_finalize_fuse_shared(
gemm2_out: torch.Tensor,
expanded_idx_to_permuted_idx: torch.Tensor,
expert_weights: torch.Tensor,
shared_output: Optional[torch.Tensor],
top_k: int,
enable_pdl: bool = False,
) -> torch.Tensor:
assert gemm2_out.dtype == torch.bfloat16
assert expert_weights.dtype in (torch.float32, torch.bfloat16)
assert expanded_idx_to_permuted_idx.dtype == torch.int32
assert gemm2_out.dim() == 2
assert expert_weights.dim() == 2
num_tokens, top_k_check = expert_weights.shape
assert top_k_check == top_k
hidden_dim = gemm2_out.shape[1]
if shared_output is not None:
assert shared_output.dtype == torch.bfloat16
assert shared_output.dim() == 2
assert shared_output.shape[0] == num_tokens
hidden_dim = shared_output.shape[1]
assert hidden_dim <= gemm2_out.shape[1]
out = torch.empty(
num_tokens, hidden_dim, dtype=torch.bfloat16, device=gemm2_out.device
)
if shared_output is None:
shared_output = gemm2_out.new_empty((0, 0), dtype=torch.bfloat16)
_jit_module().moe_finalize_fuse_shared(
out,
gemm2_out,
expanded_idx_to_permuted_idx,
expert_weights,
shared_output,
int(top_k),
bool(enable_pdl),
)
return out
+1
View File
@@ -866,6 +866,7 @@ class Envs:
# Sglang Cache Dir
SGLANG_CACHE_DIR = EnvStr(os.path.expanduser("~/.cache/sglang"))
SGLANG_FLASHINFER_AUTOTUNE_CACHE = EnvBool(True)
SGLANG_ENABLE_MOE_DEFERRED_FINALIZE = EnvBool(False)
# Plugin system
SGLANG_PLATFORM = EnvStr("")
@@ -23,6 +23,7 @@ from sglang.srt.distributed import (
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory,
)
from sglang.srt.environ import envs
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.moe import (
@@ -287,6 +288,17 @@ class FusedMoE(torch.nn.Module):
self.use_flashinfer_trtllm_moe,
self.use_deep_gemm,
)
self.supports_deferred_finalize = (
envs.SGLANG_ENABLE_MOE_DEFERRED_FINALIZE.get()
and get_moe_runner_backend().is_flashinfer_trtllm()
and isinstance(self.quant_method, ModelOptNvFp4FusedMoEMethod)
)
print_info_once(
"FlashInfer TRTLLM MoE deferred finalize is "
f"{'enabled' if self.supports_deferred_finalize else 'disabled'} "
f"(moe_runner_backend={server_args.moe_runner_backend}, "
f"quant_method={type(self.quant_method).__name__})."
)
self.quant_method.create_weights(
layer=self,
@@ -1124,6 +1136,23 @@ class FusedMoE(torch.nn.Module):
return final_hidden_states
def forward_deferred_finalize(
self, hidden_states: torch.Tensor, topk_output: TopKOutput
):
assert self.quant_method is not None
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
flashinfer_trtllm_deferred_finalize_context,
)
dispatch_output = self.dispatcher.dispatch(
hidden_states=hidden_states, topk_output=topk_output
)
with flashinfer_trtllm_deferred_finalize_context():
combine_input = self.run_moe_core(dispatch_output=dispatch_output)
return self.dispatcher.combine(combine_input=combine_input)
def run_moe_core(self, dispatch_output: DispatchOutput) -> CombineInput:
# TODO: consider using symmetric memory
return self.quant_method.apply(
@@ -1,7 +1,9 @@
from __future__ import annotations
import contextvars
from contextlib import contextmanager
from dataclasses import dataclass
from typing import TYPE_CHECKING, cast
from typing import TYPE_CHECKING, Generator, cast
import torch
from torch.nn import Module
@@ -42,6 +44,46 @@ _SGLANG_EXPERIMENTAL_LORA_OPTI = envs.SGLANG_EXPERIMENTAL_LORA_OPTI.get()
logger = __import__("logging").getLogger(__name__)
_deferred_finalize_enabled: contextvars.ContextVar[bool] = contextvars.ContextVar(
"flashinfer_trtllm_deferred_finalize_enabled", default=False
)
@dataclass
class FlashInferTrtllmDeferredFinalizeOutput:
gemm2_out: torch.Tensor
expert_weights: torch.Tensor
expanded_idx_to_permuted_idx: torch.Tensor
top_k: int
@contextmanager
def flashinfer_trtllm_deferred_finalize_context(
enabled: bool = True,
) -> Generator[None, None, None]:
token = _deferred_finalize_enabled.set(enabled)
try:
yield
finally:
_deferred_finalize_enabled.reset(token)
def finalize_flashinfer_trtllm_deferred_output(
deferred_output: FlashInferTrtllmDeferredFinalizeOutput,
shared_output: torch.Tensor,
) -> torch.Tensor:
from sglang.jit_kernel.moe_finalize_fuse_shared import moe_finalize_fuse_shared
from sglang.jit_kernel.utils import is_arch_support_pdl
return moe_finalize_fuse_shared(
deferred_output.gemm2_out,
deferred_output.expanded_idx_to_permuted_idx,
deferred_output.expert_weights,
shared_output,
deferred_output.top_k,
enable_pdl=is_arch_support_pdl(),
)
def round_up_to_multiple(x: int, m: int) -> int:
"""Round up *x* to the nearest multiple of *m*."""
@@ -922,34 +964,45 @@ def fused_experts_none_to_flashinfer_trtllm_fp4(
else:
gemm1_clamp_limit = None
num_tokens = hs_fp4.shape[0]
hidden_size = (
hs_fp4.shape[-1] * 2 if hs_fp4.dtype == torch.uint8 else hs_fp4.shape[-1]
)
_provided = _moe_output_buf.get()
_symm_required = is_allocation_symmetric()
if (
_provided is not None
and _provided.shape == (num_tokens, hidden_size)
and _provided.dtype == hidden_states.dtype
and _provided.device == hs_fp4.device
and (
not _symm_required
or not is_symmetric_memory_enabled()
or is_tensor_in_symmetric_mempool(_provided)
)
):
symm_output = _provided
else:
with use_symmetric_memory(get_tp_group(), disabled=not _symm_required):
symm_output = torch.empty(
num_tokens, hidden_size, dtype=hidden_states.dtype, device=hs_fp4.device
)
# Fall back to routed path when topk was already materialized (e.g. sigmoid routing).
if not use_routed_topk and TopKOutputChecker.format_is_standard(topk_output):
use_routed_topk = True
defer_finalize = (
_deferred_finalize_enabled.get()
and not use_routed_topk
and TopKOutputChecker.format_is_bypassed(topk_output)
)
symm_output = None
if not defer_finalize:
num_tokens = hs_fp4.shape[0]
hidden_size = (
hs_fp4.shape[-1] * 2 if hs_fp4.dtype == torch.uint8 else hs_fp4.shape[-1]
)
_provided = _moe_output_buf.get()
_symm_required = is_allocation_symmetric()
if (
_provided is not None
and _provided.shape == (num_tokens, hidden_size)
and _provided.dtype == hidden_states.dtype
and _provided.device == hs_fp4.device
and (
not _symm_required
or not is_symmetric_memory_enabled()
or is_tensor_in_symmetric_mempool(_provided)
)
):
symm_output = _provided
else:
with use_symmetric_memory(get_tp_group(), disabled=not _symm_required):
symm_output = torch.empty(
hs_fp4.shape[0],
hidden_size,
dtype=hidden_states.dtype,
device=hs_fp4.device,
)
if use_routed_topk:
assert TopKOutputChecker.format_is_standard(topk_output)
@@ -1000,7 +1053,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp4(
if topk_config.correction_bias is None
else topk_config.correction_bias.to(hidden_states.dtype)
)
result = trtllm_fp4_block_scale_moe(
moe_kwargs = dict(
routing_logits=router_logits,
routing_bias=correction_bias,
hidden_states=hs_fp4,
@@ -1031,11 +1084,31 @@ def fused_experts_none_to_flashinfer_trtllm_fp4(
if routing_method_type is not None
else RoutingMethodType.Default
),
do_finalize=True,
do_finalize=not defer_finalize,
activation_type=activation_type,
tune_max_num_tokens=next_power_of_2(hs_fp4.shape[0]),
output=symm_output,
)[0]
)
if not defer_finalize:
moe_kwargs["output"] = symm_output
result = trtllm_fp4_block_scale_moe(**moe_kwargs)
if defer_finalize:
gemm2_out, expert_weights, expanded_idx_to_permuted_idx = result[:3]
# FIXME(kpham-sgl): flashinfer sizes this buffer from routing_logits
# dtype (fp32 in DSv3 decode) but always writes bf16 weights into it.
# Reinterpret the live bf16 prefix. Fix upstream alloc to drop this,
# tracking in https://github.com/flashinfer-ai/flashinfer/issues/3595
if expert_weights.dtype == torch.float32:
n, k = expert_weights.shape
expert_weights = expert_weights.view(torch.bfloat16).view(-1, k)[:n]
result = FlashInferTrtllmDeferredFinalizeOutput(
gemm2_out=gemm2_out,
expert_weights=expert_weights,
expanded_idx_to_permuted_idx=expanded_idx_to_permuted_idx,
top_k=topk_config.top_k,
)
else:
result = result[0]
return StandardCombineInput(hidden_states=result)
+28 -7
View File
@@ -905,7 +905,18 @@ class DeepseekV2MoE(nn.Module):
expert_location_dispatch_info=dispatch_info,
**topk_kwargs,
)
final_hidden_states = self.experts(hidden_states, topk_output)
deferred_finalize = (
shared_output is not None
and not self._shared_expert_tp1
and topk_output.format == TopKOutputFormat.BYPASSED
and self.experts.supports_deferred_finalize
)
if deferred_finalize:
final_hidden_states = self.experts.forward_deferred_finalize(
hidden_states, topk_output
)
else:
final_hidden_states = self.experts(hidden_states, topk_output)
if (
not _is_cuda
and not _is_musa
@@ -916,12 +927,22 @@ class DeepseekV2MoE(nn.Module):
current_stream.wait_stream(self.alt_stream)
final_hidden_states = maybe_fuse_routed_scale_and_shared_add(
self.experts,
final_hidden_states,
None if self._shared_expert_tp1 else shared_output,
self.routed_scaling_factor,
)
if deferred_finalize:
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
finalize_flashinfer_trtllm_deferred_output,
)
final_hidden_states = finalize_flashinfer_trtllm_deferred_output(
final_hidden_states,
shared_output,
)
else:
final_hidden_states = maybe_fuse_routed_scale_and_shared_add(
self.experts,
final_hidden_states,
None if self._shared_expert_tp1 else shared_output,
self.routed_scaling_factor,
)
if self.tp_size > 1 and not should_skip_post_experts_all_reduce(
is_tp_path=True,