[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:
@@ -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
|
||||||
@@ -866,6 +866,7 @@ class Envs:
|
|||||||
# Sglang Cache Dir
|
# Sglang Cache Dir
|
||||||
SGLANG_CACHE_DIR = EnvStr(os.path.expanduser("~/.cache/sglang"))
|
SGLANG_CACHE_DIR = EnvStr(os.path.expanduser("~/.cache/sglang"))
|
||||||
SGLANG_FLASHINFER_AUTOTUNE_CACHE = EnvBool(True)
|
SGLANG_FLASHINFER_AUTOTUNE_CACHE = EnvBool(True)
|
||||||
|
SGLANG_ENABLE_MOE_DEFERRED_FINALIZE = EnvBool(False)
|
||||||
|
|
||||||
# Plugin system
|
# Plugin system
|
||||||
SGLANG_PLATFORM = EnvStr("")
|
SGLANG_PLATFORM = EnvStr("")
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from sglang.srt.distributed import (
|
|||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
use_symmetric_memory,
|
use_symmetric_memory,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
|
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.dp_attention import is_allocation_symmetric
|
||||||
from sglang.srt.layers.moe import (
|
from sglang.srt.layers.moe import (
|
||||||
@@ -287,6 +288,17 @@ class FusedMoE(torch.nn.Module):
|
|||||||
self.use_flashinfer_trtllm_moe,
|
self.use_flashinfer_trtllm_moe,
|
||||||
self.use_deep_gemm,
|
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(
|
self.quant_method.create_weights(
|
||||||
layer=self,
|
layer=self,
|
||||||
@@ -1124,6 +1136,23 @@ class FusedMoE(torch.nn.Module):
|
|||||||
|
|
||||||
return final_hidden_states
|
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:
|
def run_moe_core(self, dispatch_output: DispatchOutput) -> CombineInput:
|
||||||
# TODO: consider using symmetric memory
|
# TODO: consider using symmetric memory
|
||||||
return self.quant_method.apply(
|
return self.quant_method.apply(
|
||||||
|
|||||||
@@ -1,7 +1,9 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import contextvars
|
||||||
|
from contextlib import contextmanager
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from typing import TYPE_CHECKING, cast
|
from typing import TYPE_CHECKING, Generator, cast
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch.nn import Module
|
from torch.nn import Module
|
||||||
@@ -42,6 +44,46 @@ _SGLANG_EXPERIMENTAL_LORA_OPTI = envs.SGLANG_EXPERIMENTAL_LORA_OPTI.get()
|
|||||||
|
|
||||||
logger = __import__("logging").getLogger(__name__)
|
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:
|
def round_up_to_multiple(x: int, m: int) -> int:
|
||||||
"""Round up *x* to the nearest multiple of *m*."""
|
"""Round up *x* to the nearest multiple of *m*."""
|
||||||
@@ -922,6 +964,18 @@ def fused_experts_none_to_flashinfer_trtllm_fp4(
|
|||||||
else:
|
else:
|
||||||
gemm1_clamp_limit = None
|
gemm1_clamp_limit = None
|
||||||
|
|
||||||
|
# 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]
|
num_tokens = hs_fp4.shape[0]
|
||||||
hidden_size = (
|
hidden_size = (
|
||||||
hs_fp4.shape[-1] * 2 if hs_fp4.dtype == torch.uint8 else hs_fp4.shape[-1]
|
hs_fp4.shape[-1] * 2 if hs_fp4.dtype == torch.uint8 else hs_fp4.shape[-1]
|
||||||
@@ -943,13 +997,12 @@ def fused_experts_none_to_flashinfer_trtllm_fp4(
|
|||||||
else:
|
else:
|
||||||
with use_symmetric_memory(get_tp_group(), disabled=not _symm_required):
|
with use_symmetric_memory(get_tp_group(), disabled=not _symm_required):
|
||||||
symm_output = torch.empty(
|
symm_output = torch.empty(
|
||||||
num_tokens, hidden_size, dtype=hidden_states.dtype, device=hs_fp4.device
|
hs_fp4.shape[0],
|
||||||
|
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
|
|
||||||
|
|
||||||
if use_routed_topk:
|
if use_routed_topk:
|
||||||
assert TopKOutputChecker.format_is_standard(topk_output)
|
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
|
if topk_config.correction_bias is None
|
||||||
else topk_config.correction_bias.to(hidden_states.dtype)
|
else topk_config.correction_bias.to(hidden_states.dtype)
|
||||||
)
|
)
|
||||||
result = trtllm_fp4_block_scale_moe(
|
moe_kwargs = dict(
|
||||||
routing_logits=router_logits,
|
routing_logits=router_logits,
|
||||||
routing_bias=correction_bias,
|
routing_bias=correction_bias,
|
||||||
hidden_states=hs_fp4,
|
hidden_states=hs_fp4,
|
||||||
@@ -1031,11 +1084,31 @@ def fused_experts_none_to_flashinfer_trtllm_fp4(
|
|||||||
if routing_method_type is not None
|
if routing_method_type is not None
|
||||||
else RoutingMethodType.Default
|
else RoutingMethodType.Default
|
||||||
),
|
),
|
||||||
do_finalize=True,
|
do_finalize=not defer_finalize,
|
||||||
activation_type=activation_type,
|
activation_type=activation_type,
|
||||||
tune_max_num_tokens=next_power_of_2(hs_fp4.shape[0]),
|
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)
|
return StandardCombineInput(hidden_states=result)
|
||||||
|
|
||||||
|
|||||||
@@ -905,6 +905,17 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
expert_location_dispatch_info=dispatch_info,
|
expert_location_dispatch_info=dispatch_info,
|
||||||
**topk_kwargs,
|
**topk_kwargs,
|
||||||
)
|
)
|
||||||
|
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)
|
final_hidden_states = self.experts(hidden_states, topk_output)
|
||||||
if (
|
if (
|
||||||
not _is_cuda
|
not _is_cuda
|
||||||
@@ -916,6 +927,16 @@ class DeepseekV2MoE(nn.Module):
|
|||||||
|
|
||||||
current_stream.wait_stream(self.alt_stream)
|
current_stream.wait_stream(self.alt_stream)
|
||||||
|
|
||||||
|
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(
|
final_hidden_states = maybe_fuse_routed_scale_and_shared_add(
|
||||||
self.experts,
|
self.experts,
|
||||||
final_hidden_states,
|
final_hidden_states,
|
||||||
|
|||||||
Reference in New Issue
Block a user