Upstream: xllm/kernels/cuda/moe/moe_topk_softmax_kernels.cuh (Apache 2.0)
Adapted: CHECK→TORCH_CHECK, include path fix, cuda/functional guard, pybind11
Call chain now:
qwen3_5.py:_pure_pytorch_experts()
→ _ex_moe_topk_softmax (fused CUB kernel, 1 launch)
→ fallback: torch.softmax + torch.topk (3 launches)
Files:
ex_engine/csrc/moe/moe_topk_softmax_kernels.cuh — xllm kernel (adapted)
ex_engine/csrc/moe/device_utils.cuh — xllm device utils
ex_engine/csrc/moe/moe_topk_softmax_ext.cu — pybind11 wrapper
ex_engine/python/moe_topk.py — JIT loader (same pattern as flash_qla_sm70)
qwen3_5.py — import + use in _pure_pytorch_experts()
patch_ops.sh — deploy kernel sources for JIT
80 lines
2.6 KiB
Plaintext
80 lines
2.6 KiB
Plaintext
/* Copyright 2025 The xLLM Authors. All Rights Reserved.
|
|
|
|
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
|
|
|
|
https://github.com/jd-opensource/xllm/blob/main/LICENSE
|
|
|
|
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 <cub/cub.cuh>
|
|
|
|
namespace xllm::kernel::cuda {
|
|
|
|
#define WARP_SIZE 32
|
|
|
|
#define MAX(a, b) ((a) > (b) ? (a) : (b))
|
|
#define MIN(a, b) ((a) < (b) ? (a) : (b))
|
|
|
|
// Aligned array type
|
|
template <typename T,
|
|
// Number of elements in the array
|
|
int N,
|
|
// Alignment requirement in bytes
|
|
int Alignment = sizeof(T) * N>
|
|
class alignas(Alignment) AlignedArray {
|
|
T data[N];
|
|
};
|
|
|
|
#define XLLM_SHFL_XOR_SYNC(mask, var, lane_mask) \
|
|
__shfl_xor_sync((mask), (var), (lane_mask))
|
|
#define XLLM_SHFL_XOR_SYNC_WIDTH(mask, var, lane_mask, width) \
|
|
__shfl_xor_sync((mask), (var), (lane_mask), (width))
|
|
|
|
// Define reduction operators based on CUDA version
|
|
// CUDA 13 (12.9+) deprecated cub::Max/Min in favor of cuda::maximum/minimum
|
|
#if CUDA_VERSION >= 12090
|
|
using MaxReduceOp = ::cuda::maximum<>;
|
|
using MinReduceOp = ::cuda::minimum<>;
|
|
#else
|
|
using MaxReduceOp = cub::Max;
|
|
using MinReduceOp = cub::Min;
|
|
#endif
|
|
|
|
template <typename T>
|
|
__device__ float convert_to_float(T x) {
|
|
if constexpr (std::is_same_v<T, __half>) {
|
|
return __half2float(x);
|
|
} else if constexpr (std::is_same_v<T, __nv_bfloat16>) {
|
|
return __bfloat162float(x);
|
|
} else if constexpr (std::is_same_v<T, float>) {
|
|
return x;
|
|
} else {
|
|
return static_cast<float>(x);
|
|
}
|
|
}
|
|
|
|
// Constructs some constants needed to partition the work across threads at
|
|
// compile time.
|
|
template <typename T, int EXPERTS, int BYTES_PER_LDG>
|
|
struct TopkConstants {
|
|
static constexpr int ELTS_PER_LDG = BYTES_PER_LDG / sizeof(T);
|
|
static_assert(EXPERTS / (ELTS_PER_LDG * WARP_SIZE) == 0 ||
|
|
EXPERTS % (ELTS_PER_LDG * WARP_SIZE) == 0,
|
|
"");
|
|
static constexpr int VECs_PER_THREAD =
|
|
MAX(1, EXPERTS / (ELTS_PER_LDG * WARP_SIZE));
|
|
static constexpr int VPT = VECs_PER_THREAD * ELTS_PER_LDG;
|
|
static constexpr int THREADS_PER_ROW = EXPERTS / VPT;
|
|
static constexpr int ROWS_PER_WARP = WARP_SIZE / THREADS_PER_ROW;
|
|
};
|
|
|
|
} // namespace xllm::kernel::cuda |