Compare commits
7 Commits
main
...
4646f8fd3a
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4646f8fd3a | ||
|
|
31a1e4b7c0 | ||
|
|
5c936782b7 | ||
|
|
1dba2bbf61 | ||
|
|
0ed670a2c6 | ||
|
|
ac84e40f88 | ||
|
|
b806b15688 |
@@ -15,7 +15,7 @@ command:
|
|||||||
- -tp
|
- -tp
|
||||||
- '4'
|
- '4'
|
||||||
- --max-num-seqs
|
- --max-num-seqs
|
||||||
- '1'
|
- '2'
|
||||||
- --disable-log-requests
|
- --disable-log-requests
|
||||||
- --disable-frontend-multiprocessing
|
- --disable-frontend-multiprocessing
|
||||||
- --max-num-batched-tokens
|
- --max-num-batched-tokens
|
||||||
|
|||||||
29
qwen3_6_scripts/build_corex_moe_topk_softmax.sh
Normal file
29
qwen3_6_scripts/build_corex_moe_topk_softmax.sh
Normal file
@@ -0,0 +1,29 @@
|
|||||||
|
#!/usr/bin/env bash
|
||||||
|
set -euo pipefail
|
||||||
|
|
||||||
|
VLLM_ROOT=${1:?usage: build_corex_moe_topk_softmax.sh VLLM_ROOT}
|
||||||
|
COREX_ROOT=${COREX_ROOT:-/usr/local/corex-3.2.3}
|
||||||
|
TORCH_ROOT=${TORCH_ROOT:-${COREX_ROOT}/lib64/python3/dist-packages/torch}
|
||||||
|
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
|
||||||
|
OUTPUT=${VLLM_ROOT}/corex_moe_topk_softmax.so
|
||||||
|
|
||||||
|
"${COREX_ROOT}/bin/clang++" \
|
||||||
|
-std=c++17 -O3 -shared -fPIC \
|
||||||
|
--cuda-path="${COREX_ROOT}" --cuda-gpu-arch=ivcore10 \
|
||||||
|
--no-cuda-version-check -D_GLIBCXX_USE_CXX11_ABI=0 \
|
||||||
|
-DTORCH_EXTENSION_NAME=corex_moe_topk_softmax \
|
||||||
|
-DTORCH_API_INCLUDE_EXTENSION_H \
|
||||||
|
-I"${TORCH_ROOT}/include" \
|
||||||
|
-I"${TORCH_ROOT}/include/torch/csrc/api/include" \
|
||||||
|
-I"${TORCH_ROOT}/include/TH" -I"${TORCH_ROOT}/include/THC" \
|
||||||
|
-I/usr/local/include/python3.10 \
|
||||||
|
-I"${COREX_ROOT}/include" \
|
||||||
|
-I"${SCRIPT_DIR}" \
|
||||||
|
"${SCRIPT_DIR}/corex_moe_topk_softmax.cu" \
|
||||||
|
-L"${TORCH_ROOT}/lib" -L"${COREX_ROOT}/lib64" \
|
||||||
|
-Wl,-rpath,"${TORCH_ROOT}/lib" -Wl,-rpath,"${COREX_ROOT}/lib64" \
|
||||||
|
-ltorch_python -ltorch_cuda -ltorch_cpu -ltorch \
|
||||||
|
-lc10_cuda -lc10 -lcudart -o "${OUTPUT}"
|
||||||
|
|
||||||
|
test -s "${OUTPUT}"
|
||||||
|
printf '[ok] CoreX MoE topk+softmax extension %s\n' "${OUTPUT}"
|
||||||
44
qwen3_6_scripts/corex_moe_topk_softmax.cu
Normal file
44
qwen3_6_scripts/corex_moe_topk_softmax.cu
Normal file
@@ -0,0 +1,44 @@
|
|||||||
|
// corex_moe_topk_softmax.cu — CUB-based fused topk+softmax for MoE routing
|
||||||
|
//
|
||||||
|
// Source: upstream_ref/xllm/xllm/core/kernels/cuda/moe/moe_topk_softmax_kernels.cuh
|
||||||
|
// Adapted from vllm v0.7.3 / TensorRT-LLM v0.7.1 topk_softmax_kernels.cu
|
||||||
|
//
|
||||||
|
// Builds as corex_moe_topk_softmax.so via corex clang++ (ivcore10)
|
||||||
|
// Loaded at runtime: from vllm import corex_moe_topk_softmax
|
||||||
|
//
|
||||||
|
// Replaces: torch.softmax + torch.topk in qwen3_5.py MoE routing
|
||||||
|
// Performance: single fused kernel vs 2 separate PyTorch ops
|
||||||
|
|
||||||
|
#include "moe_topk_softmax_kernels.cuh"
|
||||||
|
#include <torch/extension.h>
|
||||||
|
|
||||||
|
using namespace xllm::kernel::cuda;
|
||||||
|
|
||||||
|
// Python-facing wrapper matching wudixzy corex_*.so convention
|
||||||
|
std::tuple<torch::Tensor, torch::Tensor> moe_topk_softmax(
|
||||||
|
torch::Tensor gating_output, // (T, num_experts)
|
||||||
|
int64_t topk,
|
||||||
|
bool renormalize) {
|
||||||
|
|
||||||
|
int64_t num_tokens = gating_output.size(0);
|
||||||
|
|
||||||
|
auto topk_weights = torch::empty(
|
||||||
|
{num_tokens, topk},
|
||||||
|
torch::dtype(torch::kFloat32).device(gating_output.device()));
|
||||||
|
auto topk_indices = torch::empty(
|
||||||
|
{num_tokens, topk},
|
||||||
|
torch::dtype(torch::kInt32).device(gating_output.device()));
|
||||||
|
|
||||||
|
topk_softmax(
|
||||||
|
topk_weights, topk_indices, gating_output,
|
||||||
|
renormalize,
|
||||||
|
/*moe_softcapping=*/0.0,
|
||||||
|
/*correction_bias=*/std::nullopt);
|
||||||
|
|
||||||
|
return std::make_tuple(topk_weights, topk_indices);
|
||||||
|
}
|
||||||
|
|
||||||
|
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||||
|
m.def("moe_topk_softmax", &moe_topk_softmax,
|
||||||
|
"CUB-based fused topk+softmax for MoE routing (xllm upstream)");
|
||||||
|
}
|
||||||
81
qwen3_6_scripts/device_utils.cuh
Normal file
81
qwen3_6_scripts/device_utils.cuh
Normal file
@@ -0,0 +1,81 @@
|
|||||||
|
/* 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
|
||||||
|
|
||||||
|
// BI-V100 corex CUB compatibility
|
||||||
|
#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
|
||||||
BIN
qwen3_6_scripts/flash_qla_sm70/build/.ninja_deps
Normal file
BIN
qwen3_6_scripts/flash_qla_sm70/build/.ninja_deps
Normal file
Binary file not shown.
5
qwen3_6_scripts/flash_qla_sm70/build/.ninja_log
Normal file
5
qwen3_6_scripts/flash_qla_sm70/build/.ninja_log
Normal file
@@ -0,0 +1,5 @@
|
|||||||
|
# ninja log v5
|
||||||
|
0 61739 1786467036068204659 gdn_forward.cuda.o 4fbd18c8f06e5181
|
||||||
|
61739 62033 1786467036388208334 flash_qla_sm70_gdn_strided.so a5d04d69a8ccfcee
|
||||||
|
0 60985 1786469746403441679 gdn_forward.cuda.o 15f5cb32976bd0b3
|
||||||
|
60985 61271 1786469746711445255 flash_qla_sm70_gdn_strided.so a5d04d69a8ccfcee
|
||||||
31
qwen3_6_scripts/flash_qla_sm70/build/build.ninja
Normal file
31
qwen3_6_scripts/flash_qla_sm70/build/build.ninja
Normal file
@@ -0,0 +1,31 @@
|
|||||||
|
ninja_required_version = 1.3
|
||||||
|
cxx = c++
|
||||||
|
nvcc = /usr/local/corex/bin/clang++
|
||||||
|
|
||||||
|
cflags = -DTORCH_EXTENSION_NAME=flash_qla_sm70_gdn_strided -DTORCH_API_INCLUDE_EXTENSION_H -DPYBIND11_COMPILER_TYPE=\"_gcc\" -DPYBIND11_STDLIB=\"_libstdcpp\" -DPYBIND11_BUILD_ABI=\"_cxxabi1011\" -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include/torch/csrc/api/include -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include/TH -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include/THC -isystem /usr/local/corex/include -isystem /usr/local/include/python3.10 -D_GLIBCXX_USE_CXX11_ABI=0 -fPIC -std=c++17 -O3
|
||||||
|
post_cflags =
|
||||||
|
cuda_cflags = -DTORCH_EXTENSION_NAME=flash_qla_sm70_gdn_strided -DTORCH_API_INCLUDE_EXTENSION_H -DPYBIND11_COMPILER_TYPE=\"_gcc\" -DPYBIND11_STDLIB=\"_libstdcpp\" -DPYBIND11_BUILD_ABI=\"_cxxabi1011\" -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include/torch/csrc/api/include -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include/TH -isystem /usr/local/corex/lib64/python3/dist-packages/torch/include/THC -isystem /usr/local/corex/include -isystem /usr/local/include/python3.10 -D_GLIBCXX_USE_CXX11_ABI=0 -D__CUDA_NO_HALF_OPERATORS__ -D__CUDA_NO_HALF_CONVERSIONS__ -D__CUDA_NO_BFLOAT16_CONVERSIONS__ -D__CUDA_NO_HALF2_OPERATORS__ -D__ILUVATAR__ -D__ILUVATAR_WORKAROUND__ -D__ILUVATAR_DIAG__ -cl-single-precision-constant -fPIC -mllvm --bonus-inst-threshold=0 -O3 --cuda-gpu-arch=ivcore10 --cuda-path=/usr/local/corex -std=c++17
|
||||||
|
cuda_post_cflags =
|
||||||
|
cuda_dlink_post_cflags =
|
||||||
|
ldflags = -shared -L/usr/local/corex/lib64/python3/dist-packages/torch/lib -lc10 -lc10_cuda -ltorch_cpu -ltorch_cuda -ltorch -ltorch_python -L/usr/local/corex/lib64 -lcudart
|
||||||
|
|
||||||
|
rule compile
|
||||||
|
command = $cxx -MMD -MF $out.d $cflags -c $in -o $out $post_cflags
|
||||||
|
depfile = $out.d
|
||||||
|
deps = gcc
|
||||||
|
|
||||||
|
rule cuda_compile
|
||||||
|
command = $nvcc $cuda_cflags -c $in -o $out $cuda_post_cflags
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
rule link
|
||||||
|
command = $cxx $in $ldflags -o $out
|
||||||
|
|
||||||
|
build gdn_forward.cuda.o: cuda_compile /workspace/qwen3_6_scripts/flash_qla_sm70/csrc/gdn_forward.cu
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
build flash_qla_sm70_gdn_strided.so: link gdn_forward.cuda.o
|
||||||
|
|
||||||
|
default flash_qla_sm70_gdn_strided.so
|
||||||
BIN
qwen3_6_scripts/flash_qla_sm70/build/flash_qla_sm70_gdn_strided.so
Executable file
BIN
qwen3_6_scripts/flash_qla_sm70/build/flash_qla_sm70_gdn_strided.so
Executable file
Binary file not shown.
BIN
qwen3_6_scripts/flash_qla_sm70/build/gdn_forward.cuda.o
Normal file
BIN
qwen3_6_scripts/flash_qla_sm70/build/gdn_forward.cuda.o
Normal file
Binary file not shown.
601
qwen3_6_scripts/moe_topk_sigmoid_kernels.cuh
Normal file
601
qwen3_6_scripts/moe_topk_sigmoid_kernels.cuh
Normal file
@@ -0,0 +1,601 @@
|
|||||||
|
// Adapt from
|
||||||
|
// https://github.com/vllm-project/vllm/blob/v0.7.3/csrc/moe/topk_softmax_kernels.cu
|
||||||
|
// which is originally adapted from
|
||||||
|
// https://github.com/NVIDIA/TensorRT-LLM/blob/v0.7.1/cpp/tensorrt_llm/kernels/mixtureOfExperts/moe_kernels.cu
|
||||||
|
/* Copyright 2025 SGLang Team. 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
|
||||||
|
|
||||||
|
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.
|
||||||
|
==============================================================================*/
|
||||||
|
|
||||||
|
#include <ATen/cuda/CUDAContext.h>
|
||||||
|
#include <c10/cuda/CUDAGuard.h>
|
||||||
|
#include <torch/all.h>
|
||||||
|
|
||||||
|
#include <cub/util_type.cuh>
|
||||||
|
|
||||||
|
#include "device_utils.cuh"
|
||||||
|
|
||||||
|
namespace {
|
||||||
|
|
||||||
|
using namespace xllm::kernel::cuda;
|
||||||
|
|
||||||
|
// ====================== Sigmoid things ===============================
|
||||||
|
// We have our own implementation of sigmoid here so we can support transposing
|
||||||
|
// the output in the sigmoid kernel when we extend this module to support
|
||||||
|
// expert-choice routing.
|
||||||
|
template <typename T, int TPB>
|
||||||
|
__launch_bounds__(TPB) __global__
|
||||||
|
void moe_sigmoid(const T* input,
|
||||||
|
const bool* finished,
|
||||||
|
float* output,
|
||||||
|
const int num_cols,
|
||||||
|
const float* correction_bias) {
|
||||||
|
const int thread_row_offset = blockIdx.x * num_cols;
|
||||||
|
|
||||||
|
// Don't touch finished rows.
|
||||||
|
if ((finished != nullptr) && finished[blockIdx.x]) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
// First pass: Apply transformation, find max, and write transformed values to
|
||||||
|
// output
|
||||||
|
for (int ii = threadIdx.x; ii < num_cols; ii += TPB) {
|
||||||
|
const int idx = thread_row_offset + ii;
|
||||||
|
float val = convert_to_float<T>(input[idx]);
|
||||||
|
|
||||||
|
val = 1.0f / (1.0f + expf(-val));
|
||||||
|
|
||||||
|
// Apply correction bias if provided
|
||||||
|
if (correction_bias != nullptr) {
|
||||||
|
val = val + correction_bias[ii];
|
||||||
|
}
|
||||||
|
|
||||||
|
output[idx] = val; // Store transformed value
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int TPB>
|
||||||
|
__launch_bounds__(TPB) __global__
|
||||||
|
void moe_topK(const float* inputs_after_sigmoid,
|
||||||
|
const bool* finished,
|
||||||
|
float* output,
|
||||||
|
int* indices,
|
||||||
|
const int num_experts,
|
||||||
|
const int k,
|
||||||
|
const int start_expert,
|
||||||
|
const int end_expert,
|
||||||
|
const bool renormalize,
|
||||||
|
const float* correction_bias) {
|
||||||
|
using cub_kvp = cub::KeyValuePair<int, float>;
|
||||||
|
using BlockReduce = cub::BlockReduce<cub_kvp, TPB>;
|
||||||
|
__shared__ typename BlockReduce::TempStorage tmpStorage;
|
||||||
|
|
||||||
|
cub_kvp thread_kvp;
|
||||||
|
cub::ArgMax arg_max;
|
||||||
|
|
||||||
|
const int block_row = blockIdx.x;
|
||||||
|
|
||||||
|
const bool row_is_active = finished ? !finished[block_row] : true;
|
||||||
|
const int thread_read_offset = blockIdx.x * num_experts;
|
||||||
|
float row_sum_for_renormalize = 0;
|
||||||
|
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||||
|
thread_kvp.key = 0;
|
||||||
|
thread_kvp.value = -1.f; // This is OK because inputs are probabilities
|
||||||
|
|
||||||
|
cub_kvp inp_kvp;
|
||||||
|
for (int expert = threadIdx.x; expert < num_experts; expert += TPB) {
|
||||||
|
const int idx = thread_read_offset + expert;
|
||||||
|
inp_kvp.key = expert;
|
||||||
|
inp_kvp.value = inputs_after_sigmoid[idx];
|
||||||
|
|
||||||
|
for (int prior_k = 0; prior_k < k_idx; ++prior_k) {
|
||||||
|
const int prior_winning_expert = indices[k * block_row + prior_k];
|
||||||
|
|
||||||
|
if (prior_winning_expert == expert) {
|
||||||
|
inp_kvp = thread_kvp;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
thread_kvp = arg_max(inp_kvp, thread_kvp);
|
||||||
|
}
|
||||||
|
|
||||||
|
const cub_kvp result_kvp =
|
||||||
|
BlockReduce(tmpStorage).Reduce(thread_kvp, arg_max);
|
||||||
|
if (threadIdx.x == 0) {
|
||||||
|
// Ignore experts the node isn't responsible for with expert parallelism
|
||||||
|
const int expert = result_kvp.key;
|
||||||
|
const bool node_uses_expert =
|
||||||
|
expert >= start_expert && expert < end_expert;
|
||||||
|
const bool should_process_row = row_is_active && node_uses_expert;
|
||||||
|
|
||||||
|
const int idx = k * block_row + k_idx;
|
||||||
|
float val = result_kvp.value;
|
||||||
|
if (correction_bias != nullptr) {
|
||||||
|
val -= correction_bias[expert];
|
||||||
|
}
|
||||||
|
output[idx] = val;
|
||||||
|
indices[idx] = should_process_row ? (expert - start_expert) : num_experts;
|
||||||
|
assert(indices[idx] >= 0);
|
||||||
|
row_sum_for_renormalize += val;
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
}
|
||||||
|
|
||||||
|
if (renormalize && threadIdx.x == 0) {
|
||||||
|
float row_sum_for_renormalize_inv = 1.f / row_sum_for_renormalize;
|
||||||
|
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||||
|
const int idx = k * block_row + k_idx;
|
||||||
|
output[idx] = output[idx] * row_sum_for_renormalize_inv;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ====================== TopK sigmoid things ===============================
|
||||||
|
|
||||||
|
/*
|
||||||
|
A Top-K gating sigmoid written to exploit when the number of experts in the
|
||||||
|
MoE layers are a small power of 2. This allows us to cleanly share the rows
|
||||||
|
among the threads in a single warp and eliminate communication between warps
|
||||||
|
(so no need to use shared mem).
|
||||||
|
|
||||||
|
It fuses the sigmoid, max and argmax into a single kernel.
|
||||||
|
|
||||||
|
Limitations:
|
||||||
|
1) This implementation is intended for when the number of experts is a small
|
||||||
|
power of 2. 2) This implementation assumes k is small, but will work for any
|
||||||
|
k.
|
||||||
|
*/
|
||||||
|
|
||||||
|
template <typename T,
|
||||||
|
int VPT,
|
||||||
|
int NUM_EXPERTS,
|
||||||
|
int WARPS_PER_CTA,
|
||||||
|
int BYTES_PER_LDG>
|
||||||
|
__launch_bounds__(WARPS_PER_CTA* WARP_SIZE) __global__
|
||||||
|
void topk_gating_sigmoid(const T* input,
|
||||||
|
const bool* finished,
|
||||||
|
float* output,
|
||||||
|
const int num_rows,
|
||||||
|
int* indices,
|
||||||
|
const int k,
|
||||||
|
const int start_expert,
|
||||||
|
const int end_expert,
|
||||||
|
const bool renormalize,
|
||||||
|
const float* correction_bias) {
|
||||||
|
// We begin by enforcing compile time assertions and setting up compile time
|
||||||
|
// constants.
|
||||||
|
static_assert(VPT == (VPT & -VPT), "VPT must be power of 2");
|
||||||
|
static_assert(NUM_EXPERTS == (NUM_EXPERTS & -NUM_EXPERTS),
|
||||||
|
"NUM_EXPERTS must be power of 2");
|
||||||
|
static_assert(BYTES_PER_LDG == (BYTES_PER_LDG & -BYTES_PER_LDG),
|
||||||
|
"BYTES_PER_LDG must be power of 2");
|
||||||
|
static_assert(BYTES_PER_LDG <= 16, "BYTES_PER_LDG must be leq 16");
|
||||||
|
|
||||||
|
// Number of bytes each thread pulls in per load
|
||||||
|
static constexpr int ELTS_PER_LDG = BYTES_PER_LDG / sizeof(T);
|
||||||
|
static constexpr int ELTS_PER_ROW = NUM_EXPERTS;
|
||||||
|
static constexpr int THREADS_PER_ROW = ELTS_PER_ROW / VPT;
|
||||||
|
static constexpr int LDG_PER_THREAD = VPT / ELTS_PER_LDG;
|
||||||
|
|
||||||
|
// Restrictions based on previous section.
|
||||||
|
static_assert(
|
||||||
|
VPT % ELTS_PER_LDG == 0,
|
||||||
|
"The elements per thread must be a multiple of the elements per ldg");
|
||||||
|
static_assert(WARP_SIZE % THREADS_PER_ROW == 0,
|
||||||
|
"The threads per row must cleanly divide the threads per warp");
|
||||||
|
static_assert(THREADS_PER_ROW == (THREADS_PER_ROW & -THREADS_PER_ROW),
|
||||||
|
"THREADS_PER_ROW must be power of 2");
|
||||||
|
static_assert(THREADS_PER_ROW <= WARP_SIZE,
|
||||||
|
"THREADS_PER_ROW can be at most warp size");
|
||||||
|
|
||||||
|
// We have NUM_EXPERTS elements per row. We specialize for small #experts
|
||||||
|
static constexpr int ELTS_PER_WARP = WARP_SIZE * VPT;
|
||||||
|
static constexpr int ROWS_PER_WARP = ELTS_PER_WARP / ELTS_PER_ROW;
|
||||||
|
static constexpr int ROWS_PER_CTA = WARPS_PER_CTA * ROWS_PER_WARP;
|
||||||
|
|
||||||
|
// Restrictions for previous section.
|
||||||
|
static_assert(ELTS_PER_WARP % ELTS_PER_ROW == 0,
|
||||||
|
"The elts per row must cleanly divide the total elt per warp");
|
||||||
|
|
||||||
|
// ===================== From this point, we finally start computing run-time
|
||||||
|
// variables. ========================
|
||||||
|
|
||||||
|
// Compute CTA and warp rows. We pack multiple rows into a single warp, and a
|
||||||
|
// block contains WARPS_PER_CTA warps. This, each block processes a chunk of
|
||||||
|
// rows. We start by computing the start row for each block.
|
||||||
|
const int cta_base_row = blockIdx.x * ROWS_PER_CTA;
|
||||||
|
|
||||||
|
// Now, using the base row per thread block, we compute the base row per warp.
|
||||||
|
const int warp_base_row = cta_base_row + threadIdx.y * ROWS_PER_WARP;
|
||||||
|
|
||||||
|
// The threads in a warp are split into sub-groups that will work on a row.
|
||||||
|
// We compute row offset for each thread sub-group
|
||||||
|
const int thread_row_in_warp = threadIdx.x / THREADS_PER_ROW;
|
||||||
|
const int thread_row = warp_base_row + thread_row_in_warp;
|
||||||
|
|
||||||
|
// Threads with indices out of bounds should early exit here.
|
||||||
|
if (thread_row >= num_rows) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
const bool row_is_active = finished ? !finished[thread_row] : true;
|
||||||
|
|
||||||
|
// We finally start setting up the read pointers for each thread. First, each
|
||||||
|
// thread jumps to the start of the row it will read.
|
||||||
|
const T* thread_row_ptr = input + thread_row * ELTS_PER_ROW;
|
||||||
|
|
||||||
|
// Now, we compute the group each thread belong to in order to determine the
|
||||||
|
// first column to start loads.
|
||||||
|
const int thread_group_idx = threadIdx.x % THREADS_PER_ROW;
|
||||||
|
const int first_elt_read_by_thread = thread_group_idx * ELTS_PER_LDG;
|
||||||
|
const T* thread_read_ptr = thread_row_ptr + first_elt_read_by_thread;
|
||||||
|
|
||||||
|
// Determine the pointer type to use to read in the data depending on the
|
||||||
|
// BYTES_PER_LDG template param. In theory, this can support all powers of 2
|
||||||
|
// up to 16. NOTE(woosuk): The original implementation uses CUTLASS aligned
|
||||||
|
// array here. We defined our own aligned array and use it here to avoid the
|
||||||
|
// dependency on CUTLASS.
|
||||||
|
using AccessType = AlignedArray<T, ELTS_PER_LDG>;
|
||||||
|
|
||||||
|
// Finally, we pull in the data from global mem
|
||||||
|
T row_chunk_temp[VPT];
|
||||||
|
AccessType* row_chunk_vec_ptr =
|
||||||
|
reinterpret_cast<AccessType*>(&row_chunk_temp);
|
||||||
|
const AccessType* vec_thread_read_ptr =
|
||||||
|
reinterpret_cast<const AccessType*>(thread_read_ptr);
|
||||||
|
#pragma unroll
|
||||||
|
// Note(Byron): interleaved loads to achieve better memory coalescing
|
||||||
|
// | thread[0] | thread[1] | thread[2] | thread[3] | thread[0] | thread[1] |
|
||||||
|
// thread[2] | thread[3] | ...
|
||||||
|
for (int ii = 0; ii < LDG_PER_THREAD; ++ii) {
|
||||||
|
row_chunk_vec_ptr[ii] = vec_thread_read_ptr[ii * THREADS_PER_ROW];
|
||||||
|
}
|
||||||
|
|
||||||
|
float row_chunk[VPT];
|
||||||
|
#pragma unroll
|
||||||
|
// Note(Byron): upcast logits to float32
|
||||||
|
for (int ii = 0; ii < VPT; ++ii) {
|
||||||
|
float val = convert_to_float<T>(row_chunk_temp[ii]);
|
||||||
|
val = 1.0f / (1.0f + expf(-val));
|
||||||
|
// Apply correction bias if provided
|
||||||
|
if (correction_bias != nullptr) {
|
||||||
|
/*
|
||||||
|
LDG is interleaved
|
||||||
|
|thread0 LDG| |thread1 LDG| |thread0 LDG| |thread1 LDG|
|
||||||
|
|--------- group0 --------| |----------group1 --------|
|
||||||
|
^ local2
|
||||||
|
*/
|
||||||
|
const int group_id = ii / ELTS_PER_LDG;
|
||||||
|
const int local_id = ii % ELTS_PER_LDG;
|
||||||
|
const int expert_idx = first_elt_read_by_thread +
|
||||||
|
group_id * THREADS_PER_ROW * ELTS_PER_LDG +
|
||||||
|
local_id;
|
||||||
|
val = val + correction_bias[expert_idx];
|
||||||
|
}
|
||||||
|
|
||||||
|
row_chunk[ii] = val;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Now, row_chunk contains the sigmoid of the row chunk. Now, I want to find
|
||||||
|
// the topk elements in each row, along with the max index.
|
||||||
|
int start_col = first_elt_read_by_thread;
|
||||||
|
static constexpr int COLS_PER_GROUP_LDG = ELTS_PER_LDG * THREADS_PER_ROW;
|
||||||
|
|
||||||
|
float row_sum_for_renormalize = 0;
|
||||||
|
|
||||||
|
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||||
|
// First, each thread does the local argmax
|
||||||
|
float max_val = row_chunk[0];
|
||||||
|
int expert = start_col;
|
||||||
|
#pragma unroll
|
||||||
|
for (int ldg = 0, col = start_col; ldg < LDG_PER_THREAD;
|
||||||
|
++ldg, col += COLS_PER_GROUP_LDG) {
|
||||||
|
#pragma unroll
|
||||||
|
for (int ii = 0; ii < ELTS_PER_LDG; ++ii) {
|
||||||
|
float val = row_chunk[ldg * ELTS_PER_LDG + ii];
|
||||||
|
|
||||||
|
// No check on the experts here since columns with the smallest index
|
||||||
|
// are processed first and only updated if > (not >=)
|
||||||
|
if (val > max_val) {
|
||||||
|
max_val = val;
|
||||||
|
expert = col + ii;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Now, we perform the argmax reduce. We use the butterfly pattern so threads
|
||||||
|
// reach consensus about the max. This will be useful for K > 1 so that the
|
||||||
|
// threads can agree on "who" had the max value. That thread can then blank out
|
||||||
|
// their max with -inf and the warp can run more iterations...
|
||||||
|
#pragma unroll
|
||||||
|
for (int mask = THREADS_PER_ROW / 2; mask > 0; mask /= 2) {
|
||||||
|
float other_max =
|
||||||
|
XLLM_SHFL_XOR_SYNC_WIDTH(0xffffffff, max_val, mask, THREADS_PER_ROW);
|
||||||
|
int other_expert =
|
||||||
|
XLLM_SHFL_XOR_SYNC_WIDTH(0xffffffff, expert, mask, THREADS_PER_ROW);
|
||||||
|
|
||||||
|
// We want lower indices to "win" in every thread so we break ties this
|
||||||
|
// way
|
||||||
|
if (other_max > max_val ||
|
||||||
|
(other_max == max_val && other_expert < expert)) {
|
||||||
|
max_val = other_max;
|
||||||
|
expert = other_expert;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write the max for this k iteration to global memory.
|
||||||
|
if (thread_group_idx == 0) {
|
||||||
|
// Add a guard to ignore experts not included by this node
|
||||||
|
const bool node_uses_expert =
|
||||||
|
expert >= start_expert && expert < end_expert;
|
||||||
|
const bool should_process_row = row_is_active && node_uses_expert;
|
||||||
|
|
||||||
|
// The lead thread from each sub-group will write out the final results to
|
||||||
|
// global memory. (This will be a single) thread per row of the
|
||||||
|
// input/output matrices.
|
||||||
|
const int idx = k * thread_row + k_idx;
|
||||||
|
if (correction_bias != nullptr) {
|
||||||
|
max_val -= correction_bias[expert];
|
||||||
|
}
|
||||||
|
output[idx] = max_val;
|
||||||
|
indices[idx] = should_process_row ? (expert - start_expert) : NUM_EXPERTS;
|
||||||
|
row_sum_for_renormalize += max_val;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Finally, we clear the value in the thread with the current max if there
|
||||||
|
// is another iteration to run.
|
||||||
|
if (k_idx + 1 < k) {
|
||||||
|
const int ldg_group_for_expert = expert / COLS_PER_GROUP_LDG;
|
||||||
|
const int thread_to_clear_in_group =
|
||||||
|
(expert / ELTS_PER_LDG) % THREADS_PER_ROW;
|
||||||
|
|
||||||
|
// Only the thread in the group which produced the max will reset the
|
||||||
|
// "winning" value to -inf.
|
||||||
|
if (thread_group_idx == thread_to_clear_in_group) {
|
||||||
|
const int offset_for_expert = expert % ELTS_PER_LDG;
|
||||||
|
// Safe to set to any negative value since row_chunk values must be
|
||||||
|
// between 0 and 1.
|
||||||
|
row_chunk[ldg_group_for_expert * ELTS_PER_LDG + offset_for_expert] =
|
||||||
|
-10000.f;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fuse renormalization of topk_weights into this kernel
|
||||||
|
if (renormalize && thread_group_idx == 0) {
|
||||||
|
float row_sum_for_renormalize_inv = 1.f / row_sum_for_renormalize;
|
||||||
|
#pragma unroll
|
||||||
|
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||||
|
const int idx = k * thread_row + k_idx;
|
||||||
|
output[idx] = output[idx] * row_sum_for_renormalize_inv;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
template <typename T, int EXPERTS, int WARPS_PER_TB>
|
||||||
|
void topk_gating_sigmoid_launcher_helper(const T* input,
|
||||||
|
const bool* finished,
|
||||||
|
float* output,
|
||||||
|
int* indices,
|
||||||
|
const int num_rows,
|
||||||
|
const int k,
|
||||||
|
const int start_expert,
|
||||||
|
const int end_expert,
|
||||||
|
const bool renormalize,
|
||||||
|
const float* correction_bias,
|
||||||
|
cudaStream_t stream) {
|
||||||
|
static constexpr std::size_t MAX_BYTES_PER_LDG = 16;
|
||||||
|
|
||||||
|
static constexpr int BYTES_PER_LDG =
|
||||||
|
MIN(MAX_BYTES_PER_LDG, sizeof(T) * EXPERTS);
|
||||||
|
using Constants = TopkConstants<T, EXPERTS, BYTES_PER_LDG>;
|
||||||
|
static constexpr int VPT = Constants::VPT;
|
||||||
|
static constexpr int ROWS_PER_WARP = Constants::ROWS_PER_WARP;
|
||||||
|
const int num_warps = (num_rows + ROWS_PER_WARP - 1) / ROWS_PER_WARP;
|
||||||
|
const int num_blocks = (num_warps + WARPS_PER_TB - 1) / WARPS_PER_TB;
|
||||||
|
|
||||||
|
dim3 block_dim(WARP_SIZE, WARPS_PER_TB);
|
||||||
|
topk_gating_sigmoid<T, VPT, EXPERTS, WARPS_PER_TB, BYTES_PER_LDG>
|
||||||
|
<<<num_blocks, block_dim, 0, stream>>>(input,
|
||||||
|
finished,
|
||||||
|
output,
|
||||||
|
num_rows,
|
||||||
|
indices,
|
||||||
|
k,
|
||||||
|
start_expert,
|
||||||
|
end_expert,
|
||||||
|
renormalize,
|
||||||
|
correction_bias);
|
||||||
|
}
|
||||||
|
|
||||||
|
#define LAUNCH_SIGMOID(TYPE, NUM_EXPERTS, WARPS_PER_TB) \
|
||||||
|
topk_gating_sigmoid_launcher_helper<TYPE, NUM_EXPERTS, WARPS_PER_TB>( \
|
||||||
|
gating_output, \
|
||||||
|
nullptr, \
|
||||||
|
topk_weights, \
|
||||||
|
topk_indices, \
|
||||||
|
num_tokens, \
|
||||||
|
topk, \
|
||||||
|
0, \
|
||||||
|
num_experts, \
|
||||||
|
renormalize, \
|
||||||
|
correction_bias, \
|
||||||
|
stream);
|
||||||
|
|
||||||
|
template <typename T>
|
||||||
|
void topk_gating_sigmoid_kernel_launcher(const T* gating_output,
|
||||||
|
float* topk_weights,
|
||||||
|
int* topk_indices,
|
||||||
|
float* sigmoid_workspace,
|
||||||
|
const int num_tokens,
|
||||||
|
const int num_experts,
|
||||||
|
const int topk,
|
||||||
|
const bool renormalize,
|
||||||
|
const float* correction_bias,
|
||||||
|
cudaStream_t stream) {
|
||||||
|
static constexpr int WARPS_PER_TB = 4;
|
||||||
|
switch (num_experts) {
|
||||||
|
case 1:
|
||||||
|
LAUNCH_SIGMOID(T, 1, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 2:
|
||||||
|
LAUNCH_SIGMOID(T, 2, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 4:
|
||||||
|
LAUNCH_SIGMOID(T, 4, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 8:
|
||||||
|
LAUNCH_SIGMOID(T, 8, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 16:
|
||||||
|
LAUNCH_SIGMOID(T, 16, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 32:
|
||||||
|
LAUNCH_SIGMOID(T, 32, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 64:
|
||||||
|
LAUNCH_SIGMOID(T, 64, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 128:
|
||||||
|
LAUNCH_SIGMOID(T, 128, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 256:
|
||||||
|
LAUNCH_SIGMOID(T, 256, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
default: {
|
||||||
|
TORCH_CHECK(sigmoid_workspace != nullptr,
|
||||||
|
"sigmoid_workspace must be provided for num_experts that are "
|
||||||
|
"not a power of 2.");
|
||||||
|
static constexpr int TPB = 256;
|
||||||
|
moe_sigmoid<T, TPB><<<num_tokens, TPB, 0, stream>>>(gating_output,
|
||||||
|
nullptr,
|
||||||
|
sigmoid_workspace,
|
||||||
|
num_experts,
|
||||||
|
correction_bias);
|
||||||
|
moe_topK<TPB><<<num_tokens, TPB, 0, stream>>>(sigmoid_workspace,
|
||||||
|
nullptr,
|
||||||
|
topk_weights,
|
||||||
|
topk_indices,
|
||||||
|
num_experts,
|
||||||
|
topk,
|
||||||
|
0,
|
||||||
|
num_experts,
|
||||||
|
renormalize,
|
||||||
|
correction_bias);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} // namespace
|
||||||
|
|
||||||
|
namespace xllm::kernel::cuda {
|
||||||
|
void topk_sigmoid(torch::Tensor& topk_weights, // [num_tokens, topk]
|
||||||
|
torch::Tensor& topk_indices, // [num_tokens, topk]
|
||||||
|
torch::Tensor& gating_output, // [num_tokens, num_experts]
|
||||||
|
const bool renormalize,
|
||||||
|
const std::optional<torch::Tensor>& correction_bias) {
|
||||||
|
// Check data type
|
||||||
|
CHECK(gating_output.scalar_type() == at::ScalarType::Float ||
|
||||||
|
gating_output.scalar_type() == at::ScalarType::Half ||
|
||||||
|
gating_output.scalar_type() == at::ScalarType::BFloat16)
|
||||||
|
<< "gating_output must be float32, float16, or bfloat16";
|
||||||
|
|
||||||
|
// Check dimensions
|
||||||
|
CHECK(gating_output.dim() == 2)
|
||||||
|
<< "gating_output must be 2D tensor [num_tokens, num_experts]";
|
||||||
|
CHECK(topk_weights.dim() == 2)
|
||||||
|
<< "topk_weights must be 2D tensor [num_tokens, topk]";
|
||||||
|
CHECK(topk_indices.dim() == 2)
|
||||||
|
<< "topk_indices must be 2D tensor [num_tokens, topk]";
|
||||||
|
|
||||||
|
// Check shapes
|
||||||
|
CHECK(gating_output.size(0) == topk_weights.size(0))
|
||||||
|
<< "First dimension of topk_weights must match num_tokens in "
|
||||||
|
"gating_output";
|
||||||
|
CHECK(gating_output.size(0) == topk_indices.size(0))
|
||||||
|
<< "First dimension of topk_indices must match num_tokens in "
|
||||||
|
"gating_output";
|
||||||
|
CHECK(topk_weights.size(-1) == topk_indices.size(-1))
|
||||||
|
<< "Second dimension of topk_indices must match topk in topk_weights";
|
||||||
|
CHECK(topk_weights.size(-1) <= gating_output.size(-1))
|
||||||
|
<< "topk must be less than or equal to num_experts";
|
||||||
|
|
||||||
|
const int num_experts = static_cast<int>(gating_output.size(-1));
|
||||||
|
const int num_tokens = static_cast<int>(gating_output.size(0));
|
||||||
|
const int topk = static_cast<int>(topk_weights.size(-1));
|
||||||
|
|
||||||
|
const bool is_pow_2 =
|
||||||
|
(num_experts != 0) && ((num_experts & (num_experts - 1)) == 0);
|
||||||
|
const bool needs_workspace = !is_pow_2 || num_experts > 256;
|
||||||
|
const int64_t workspace_size = needs_workspace ? num_tokens * num_experts : 0;
|
||||||
|
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(gating_output));
|
||||||
|
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||||
|
torch::Tensor sigmoid_workspace = torch::empty(
|
||||||
|
{workspace_size}, gating_output.options().dtype(at::ScalarType::Float));
|
||||||
|
|
||||||
|
const at::ScalarType dtype = gating_output.scalar_type();
|
||||||
|
|
||||||
|
// Validate correction_bias if provided - must always be float32
|
||||||
|
const float* bias_ptr = nullptr;
|
||||||
|
if (correction_bias.has_value()) {
|
||||||
|
const torch::Tensor& bias_tensor = correction_bias.value();
|
||||||
|
CHECK(bias_tensor.dim() == 1)
|
||||||
|
<< "correction_bias must be 1D tensor [num_experts]";
|
||||||
|
CHECK(bias_tensor.size(0) == num_experts)
|
||||||
|
<< "correction_bias size must match num_experts";
|
||||||
|
CHECK(bias_tensor.scalar_type() == at::ScalarType::Float)
|
||||||
|
<< "correction_bias must be float32, got " << bias_tensor.scalar_type();
|
||||||
|
bias_ptr = bias_tensor.data_ptr<float>();
|
||||||
|
}
|
||||||
|
|
||||||
|
if (dtype == at::ScalarType::Float) {
|
||||||
|
topk_gating_sigmoid_kernel_launcher<float>(
|
||||||
|
gating_output.data_ptr<float>(),
|
||||||
|
topk_weights.data_ptr<float>(),
|
||||||
|
topk_indices.data_ptr<int>(),
|
||||||
|
sigmoid_workspace.data_ptr<float>(),
|
||||||
|
num_tokens,
|
||||||
|
num_experts,
|
||||||
|
topk,
|
||||||
|
renormalize,
|
||||||
|
bias_ptr,
|
||||||
|
stream);
|
||||||
|
} else if (dtype == at::ScalarType::Half) {
|
||||||
|
topk_gating_sigmoid_kernel_launcher<__half>(
|
||||||
|
reinterpret_cast<const __half*>(gating_output.data_ptr<at::Half>()),
|
||||||
|
topk_weights.data_ptr<float>(),
|
||||||
|
topk_indices.data_ptr<int>(),
|
||||||
|
sigmoid_workspace.data_ptr<float>(),
|
||||||
|
num_tokens,
|
||||||
|
num_experts,
|
||||||
|
topk,
|
||||||
|
renormalize,
|
||||||
|
bias_ptr,
|
||||||
|
stream);
|
||||||
|
} else if (dtype == at::ScalarType::BFloat16) {
|
||||||
|
topk_gating_sigmoid_kernel_launcher<__nv_bfloat16>(
|
||||||
|
reinterpret_cast<const __nv_bfloat16*>(
|
||||||
|
gating_output.data_ptr<at::BFloat16>()),
|
||||||
|
topk_weights.data_ptr<float>(),
|
||||||
|
topk_indices.data_ptr<int>(),
|
||||||
|
sigmoid_workspace.data_ptr<float>(),
|
||||||
|
num_tokens,
|
||||||
|
num_experts,
|
||||||
|
topk,
|
||||||
|
renormalize,
|
||||||
|
bias_ptr,
|
||||||
|
stream);
|
||||||
|
} else {
|
||||||
|
LOG(FATAL) << "Unsupported gating_output dtype: " << dtype;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} // namespace xllm::kernel::cuda
|
||||||
840
qwen3_6_scripts/moe_topk_softmax_kernels.cuh
Normal file
840
qwen3_6_scripts/moe_topk_softmax_kernels.cuh
Normal file
@@ -0,0 +1,840 @@
|
|||||||
|
// Adapt from
|
||||||
|
// https://github.com/vllm-project/vllm/blob/v0.7.3/csrc/moe/topk_softmax_kernels.cu
|
||||||
|
// which is originally adapted from
|
||||||
|
// https://github.com/NVIDIA/TensorRT-LLM/blob/v0.7.1/cpp/tensorrt_llm/kernels/mixtureOfExperts/moe_kernels.cu
|
||||||
|
/* Copyright 2025 SGLang Team. 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
|
||||||
|
|
||||||
|
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.
|
||||||
|
==============================================================================*/
|
||||||
|
|
||||||
|
#include <ATen/cuda/CUDAContext.h>
|
||||||
|
#include <c10/cuda/CUDAGuard.h>
|
||||||
|
#include <torch/all.h>
|
||||||
|
|
||||||
|
#include <cub/util_type.cuh>
|
||||||
|
|
||||||
|
#include "device_utils.cuh"
|
||||||
|
|
||||||
|
using cub_kvp = cub::KeyValuePair<int, float>;
|
||||||
|
|
||||||
|
namespace {
|
||||||
|
|
||||||
|
using namespace xllm::kernel::cuda;
|
||||||
|
|
||||||
|
// ====================== Softmax things ===============================
|
||||||
|
// We have our own implementation of softmax here so we can support transposing
|
||||||
|
// the output in the softmax kernel when we extend this module to support
|
||||||
|
// expert-choice routing.
|
||||||
|
template <typename T, int TPB>
|
||||||
|
__launch_bounds__(TPB) __global__
|
||||||
|
void moe_softmax(const T* input,
|
||||||
|
const bool* finished,
|
||||||
|
float* output,
|
||||||
|
const int num_cols,
|
||||||
|
const float moe_softcapping,
|
||||||
|
const float* correction_bias) {
|
||||||
|
using BlockReduce = cub::BlockReduce<float, TPB>;
|
||||||
|
__shared__ typename BlockReduce::TempStorage tmpStorage;
|
||||||
|
|
||||||
|
__shared__ float normalizing_factor;
|
||||||
|
__shared__ float float_max;
|
||||||
|
|
||||||
|
const int thread_row_offset = blockIdx.x * num_cols;
|
||||||
|
|
||||||
|
float threadData(-FLT_MAX);
|
||||||
|
|
||||||
|
// Don't touch finished rows.
|
||||||
|
if ((finished != nullptr) && finished[blockIdx.x]) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
// First pass: Apply transformation, find max, and write transformed values to
|
||||||
|
// output
|
||||||
|
for (int ii = threadIdx.x; ii < num_cols; ii += TPB) {
|
||||||
|
const int idx = thread_row_offset + ii;
|
||||||
|
float val = convert_to_float<T>(input[idx]);
|
||||||
|
|
||||||
|
// Apply tanh softcapping if enabled
|
||||||
|
if (moe_softcapping != 0.0f) {
|
||||||
|
val = tanhf(val / moe_softcapping) * moe_softcapping;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply correction bias if provided
|
||||||
|
if (correction_bias != nullptr) {
|
||||||
|
val = val + correction_bias[ii];
|
||||||
|
}
|
||||||
|
|
||||||
|
output[idx] = val; // Store transformed value
|
||||||
|
threadData = max(val, threadData);
|
||||||
|
}
|
||||||
|
|
||||||
|
const float maxElem =
|
||||||
|
BlockReduce(tmpStorage).Reduce(threadData, MaxReduceOp());
|
||||||
|
|
||||||
|
if (threadIdx.x == 0) {
|
||||||
|
float_max = maxElem;
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
|
||||||
|
// Second pass: Compute sum using transformed values from output
|
||||||
|
threadData = 0;
|
||||||
|
for (int ii = threadIdx.x; ii < num_cols; ii += TPB) {
|
||||||
|
const int idx = thread_row_offset + ii;
|
||||||
|
threadData += exp((output[idx] - float_max));
|
||||||
|
}
|
||||||
|
|
||||||
|
const auto Z = BlockReduce(tmpStorage).Sum(threadData);
|
||||||
|
|
||||||
|
if (threadIdx.x == 0) {
|
||||||
|
normalizing_factor = 1.f / Z;
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
|
||||||
|
// Third pass: Compute final softmax using transformed values from output
|
||||||
|
for (int ii = threadIdx.x; ii < num_cols; ii += TPB) {
|
||||||
|
const int idx = thread_row_offset + ii;
|
||||||
|
const float softmax_val =
|
||||||
|
exp((output[idx] - float_max)) * normalizing_factor;
|
||||||
|
output[idx] = softmax_val;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
namespace moe {
|
||||||
|
struct TopKPair {
|
||||||
|
static const int PAIR = 2;
|
||||||
|
static const int MAX_INDEX = 0;
|
||||||
|
cub_kvp max;
|
||||||
|
cub_kvp secondMax;
|
||||||
|
|
||||||
|
__device__ TopKPair() {}
|
||||||
|
__device__ TopKPair(cub_kvp max, cub_kvp secondMax)
|
||||||
|
: max(max), secondMax(secondMax) {}
|
||||||
|
};
|
||||||
|
|
||||||
|
struct TopKPairArgMax {
|
||||||
|
__device__ TopKPairArgMax() {}
|
||||||
|
__device__ __forceinline__ TopKPair
|
||||||
|
operator()(const TopKPair& candidate1, const TopKPair& candidate2) const {
|
||||||
|
cub_kvp globalMax, globalSecondMax;
|
||||||
|
|
||||||
|
// Determine the global maximum
|
||||||
|
if (candidate1.max.value > candidate2.max.value) {
|
||||||
|
globalMax = candidate1.max;
|
||||||
|
} else {
|
||||||
|
globalMax = candidate2.max;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Determine the global second maximum
|
||||||
|
if (globalMax.key == candidate1.max.key) {
|
||||||
|
// If candidate1 contributed the max, compare its secondMax with
|
||||||
|
// candidate2's max
|
||||||
|
globalSecondMax = (candidate1.secondMax.value > candidate2.max.value)
|
||||||
|
? candidate1.secondMax
|
||||||
|
: candidate2.max;
|
||||||
|
} else {
|
||||||
|
// If candidate2 contributed the max, compare its secondMax with
|
||||||
|
// candidate1's max
|
||||||
|
globalSecondMax = (candidate2.secondMax.value > candidate1.max.value)
|
||||||
|
? candidate2.secondMax
|
||||||
|
: candidate1.max;
|
||||||
|
}
|
||||||
|
return TopKPair(globalMax, globalSecondMax);
|
||||||
|
}
|
||||||
|
};
|
||||||
|
} // namespace moe
|
||||||
|
|
||||||
|
template <int TPB>
|
||||||
|
__launch_bounds__(TPB) __global__
|
||||||
|
void moe_topk_fast(float* inputs_after_softmax,
|
||||||
|
const bool* finished,
|
||||||
|
float* output,
|
||||||
|
int* indices,
|
||||||
|
const int num_experts,
|
||||||
|
const int k,
|
||||||
|
const int start_expert,
|
||||||
|
const int end_expert,
|
||||||
|
const bool renormalize) {
|
||||||
|
using namespace moe;
|
||||||
|
using BlockReduce = cub::BlockReduce<TopKPair, TPB>;
|
||||||
|
__shared__ typename BlockReduce::TempStorage tmpStorage;
|
||||||
|
TopKPair thread_pair;
|
||||||
|
|
||||||
|
const int block_row = blockIdx.x;
|
||||||
|
|
||||||
|
const bool row_is_active = finished ? !finished[block_row] : true;
|
||||||
|
const int thread_read_offset = blockIdx.x * num_experts;
|
||||||
|
float row_sum_for_renormalize = 0;
|
||||||
|
// Each loop finds the top 2 elements,
|
||||||
|
// thus requiring only ⌈k/2⌉ loops (calculated as (k + 1) / 2).
|
||||||
|
for (int k_idx = 0; k_idx < (k + TopKPair::PAIR - 1) / TopKPair::PAIR;
|
||||||
|
++k_idx) {
|
||||||
|
// Initializing the top 2 elements by the minimum value.
|
||||||
|
thread_pair.max.key = 0;
|
||||||
|
thread_pair.max.value = -1.f;
|
||||||
|
thread_pair.secondMax.key = 0;
|
||||||
|
thread_pair.secondMax.value = -1.f;
|
||||||
|
|
||||||
|
cub_kvp inp_kvp;
|
||||||
|
for (int expert = threadIdx.x; expert < num_experts; expert += TPB) {
|
||||||
|
const int idx = thread_read_offset + expert;
|
||||||
|
inp_kvp.key = expert;
|
||||||
|
inp_kvp.value = inputs_after_softmax[idx];
|
||||||
|
// updating the thread_pair according to inp_kvp's value
|
||||||
|
if (inp_kvp.value > thread_pair.max.value) {
|
||||||
|
thread_pair.secondMax = thread_pair.max;
|
||||||
|
thread_pair.max = inp_kvp;
|
||||||
|
} else if (inp_kvp.value > thread_pair.secondMax.value) {
|
||||||
|
thread_pair.secondMax = inp_kvp;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
TopKPairArgMax reducer;
|
||||||
|
const TopKPair result_pair =
|
||||||
|
BlockReduce(tmpStorage).Reduce(thread_pair, reducer);
|
||||||
|
if (threadIdx.x == 0) {
|
||||||
|
#pragma unroll
|
||||||
|
// updating 2 elements to the result.
|
||||||
|
for (int i = 0; i < TopKPair::PAIR; i++) {
|
||||||
|
if (k_idx * 2 + i >= k) break;
|
||||||
|
cub_kvp result = (i == TopKPair::MAX_INDEX) ? result_pair.max
|
||||||
|
: result_pair.secondMax;
|
||||||
|
int expert = result.key;
|
||||||
|
bool node_uses_expert = expert >= start_expert && expert < end_expert;
|
||||||
|
bool should_process_row = row_is_active && node_uses_expert;
|
||||||
|
// The inputs_after_softmax is modified in-place to avoid unnecessary
|
||||||
|
// loops for finding the top k-1 value. 1.f represents the minimum
|
||||||
|
// value.
|
||||||
|
inputs_after_softmax[thread_read_offset + expert] = -1.f;
|
||||||
|
int idx = k * block_row + k_idx * 2 + i;
|
||||||
|
output[idx] = result.value;
|
||||||
|
indices[idx] =
|
||||||
|
should_process_row ? (expert - start_expert) : num_experts;
|
||||||
|
assert(indices[idx] >= 0);
|
||||||
|
row_sum_for_renormalize += result.value;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
}
|
||||||
|
|
||||||
|
if (renormalize && threadIdx.x == 0) {
|
||||||
|
float row_sum_for_renormalize_inv = 1.f / row_sum_for_renormalize;
|
||||||
|
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||||
|
const int idx = k * block_row + k_idx;
|
||||||
|
output[idx] = output[idx] * row_sum_for_renormalize_inv;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
template <int TPB>
|
||||||
|
__launch_bounds__(TPB) __global__ void moe_topK(float* inputs_after_softmax,
|
||||||
|
const bool* finished,
|
||||||
|
float* output,
|
||||||
|
int* indices,
|
||||||
|
const int num_experts,
|
||||||
|
const int k,
|
||||||
|
const int start_expert,
|
||||||
|
const int end_expert,
|
||||||
|
const bool renormalize) {
|
||||||
|
using cub_kvp = cub::KeyValuePair<int, float>;
|
||||||
|
using BlockReduce = cub::BlockReduce<cub_kvp, TPB>;
|
||||||
|
__shared__ typename BlockReduce::TempStorage tmpStorage;
|
||||||
|
|
||||||
|
cub_kvp thread_kvp;
|
||||||
|
cub::ArgMax arg_max;
|
||||||
|
|
||||||
|
const int block_row = blockIdx.x;
|
||||||
|
|
||||||
|
const bool row_is_active = finished ? !finished[block_row] : true;
|
||||||
|
const int thread_read_offset = blockIdx.x * num_experts;
|
||||||
|
float row_sum_for_renormalize = 0;
|
||||||
|
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||||
|
thread_kvp.key = 0;
|
||||||
|
thread_kvp.value = -1.f; // This is OK because inputs are probabilities
|
||||||
|
|
||||||
|
cub_kvp inp_kvp;
|
||||||
|
for (int expert = threadIdx.x; expert < num_experts; expert += TPB) {
|
||||||
|
const int idx = thread_read_offset + expert;
|
||||||
|
inp_kvp.key = expert;
|
||||||
|
inp_kvp.value = inputs_after_softmax[idx];
|
||||||
|
thread_kvp = arg_max(inp_kvp, thread_kvp);
|
||||||
|
}
|
||||||
|
|
||||||
|
const cub_kvp result_kvp =
|
||||||
|
BlockReduce(tmpStorage).Reduce(thread_kvp, arg_max);
|
||||||
|
if (threadIdx.x == 0) {
|
||||||
|
// Ignore experts the node isn't responsible for with expert parallelism
|
||||||
|
const int expert = result_kvp.key;
|
||||||
|
const bool node_uses_expert =
|
||||||
|
expert >= start_expert && expert < end_expert;
|
||||||
|
const bool should_process_row = row_is_active && node_uses_expert;
|
||||||
|
|
||||||
|
const int idx = k * block_row + k_idx;
|
||||||
|
output[idx] = result_kvp.value;
|
||||||
|
indices[idx] = should_process_row ? (expert - start_expert) : num_experts;
|
||||||
|
assert(indices[idx] >= 0);
|
||||||
|
row_sum_for_renormalize += result_kvp.value;
|
||||||
|
// The inputs_after_softmax is modified in-place to avoid unnecessary
|
||||||
|
// loops for finding the top k-1 value. 1.f represents the minimum value.
|
||||||
|
inputs_after_softmax[thread_read_offset + expert] = -1.f;
|
||||||
|
}
|
||||||
|
__syncthreads();
|
||||||
|
}
|
||||||
|
|
||||||
|
if (renormalize && threadIdx.x == 0) {
|
||||||
|
float row_sum_for_renormalize_inv = 1.f / row_sum_for_renormalize;
|
||||||
|
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||||
|
const int idx = k * block_row + k_idx;
|
||||||
|
output[idx] = output[idx] * row_sum_for_renormalize_inv;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ====================== TopK softmax things ===============================
|
||||||
|
|
||||||
|
/*
|
||||||
|
A Top-K gating softmax written to exploit when the number of experts in the
|
||||||
|
MoE layers are a small power of 2. This allows us to cleanly share the rows
|
||||||
|
among the threads in a single warp and eliminate communication between warps
|
||||||
|
(so no need to use shared mem).
|
||||||
|
|
||||||
|
It fuses the softmax, max and argmax into a single kernel.
|
||||||
|
|
||||||
|
Limitations:
|
||||||
|
1) This implementation is intended for when the number of experts is a small
|
||||||
|
power of 2. 2) This implementation assumes k is small, but will work for any
|
||||||
|
k.
|
||||||
|
*/
|
||||||
|
|
||||||
|
template <typename T,
|
||||||
|
int VPT,
|
||||||
|
int NUM_EXPERTS,
|
||||||
|
int WARPS_PER_CTA,
|
||||||
|
int BYTES_PER_LDG>
|
||||||
|
__launch_bounds__(WARPS_PER_CTA* WARP_SIZE) __global__
|
||||||
|
void topk_gating_softmax(const T* input,
|
||||||
|
const bool* finished,
|
||||||
|
float* output,
|
||||||
|
const int num_rows,
|
||||||
|
int* indices,
|
||||||
|
const int k,
|
||||||
|
const int start_expert,
|
||||||
|
const int end_expert,
|
||||||
|
const bool renormalize,
|
||||||
|
const float moe_softcapping,
|
||||||
|
const float* correction_bias) {
|
||||||
|
// We begin by enforcing compile time assertions and setting up compile time
|
||||||
|
// constants.
|
||||||
|
static_assert(VPT == (VPT & -VPT), "VPT must be power of 2");
|
||||||
|
static_assert(NUM_EXPERTS == (NUM_EXPERTS & -NUM_EXPERTS),
|
||||||
|
"NUM_EXPERTS must be power of 2");
|
||||||
|
static_assert(BYTES_PER_LDG == (BYTES_PER_LDG & -BYTES_PER_LDG),
|
||||||
|
"BYTES_PER_LDG must be power of 2");
|
||||||
|
static_assert(BYTES_PER_LDG <= 16, "BYTES_PER_LDG must be leq 16");
|
||||||
|
|
||||||
|
// Number of bytes each thread pulls in per load
|
||||||
|
static constexpr int ELTS_PER_LDG = BYTES_PER_LDG / sizeof(T);
|
||||||
|
static constexpr int ELTS_PER_ROW = NUM_EXPERTS;
|
||||||
|
static constexpr int THREADS_PER_ROW = ELTS_PER_ROW / VPT;
|
||||||
|
static constexpr int LDG_PER_THREAD = VPT / ELTS_PER_LDG;
|
||||||
|
|
||||||
|
// Restrictions based on previous section.
|
||||||
|
static_assert(
|
||||||
|
VPT % ELTS_PER_LDG == 0,
|
||||||
|
"The elements per thread must be a multiple of the elements per ldg");
|
||||||
|
static_assert(WARP_SIZE % THREADS_PER_ROW == 0,
|
||||||
|
"The threads per row must cleanly divide the threads per warp");
|
||||||
|
static_assert(THREADS_PER_ROW == (THREADS_PER_ROW & -THREADS_PER_ROW),
|
||||||
|
"THREADS_PER_ROW must be power of 2");
|
||||||
|
static_assert(THREADS_PER_ROW <= WARP_SIZE,
|
||||||
|
"THREADS_PER_ROW can be at most warp size");
|
||||||
|
|
||||||
|
// We have NUM_EXPERTS elements per row. We specialize for small #experts
|
||||||
|
static constexpr int ELTS_PER_WARP = WARP_SIZE * VPT;
|
||||||
|
static constexpr int ROWS_PER_WARP = ELTS_PER_WARP / ELTS_PER_ROW;
|
||||||
|
static constexpr int ROWS_PER_CTA = WARPS_PER_CTA * ROWS_PER_WARP;
|
||||||
|
|
||||||
|
// Restrictions for previous section.
|
||||||
|
static_assert(ELTS_PER_WARP % ELTS_PER_ROW == 0,
|
||||||
|
"The elts per row must cleanly divide the total elt per warp");
|
||||||
|
|
||||||
|
// ===================== From this point, we finally start computing run-time
|
||||||
|
// variables. ========================
|
||||||
|
|
||||||
|
// Compute CTA and warp rows. We pack multiple rows into a single warp, and a
|
||||||
|
// block contains WARPS_PER_CTA warps. This, each block processes a chunk of
|
||||||
|
// rows. We start by computing the start row for each block.
|
||||||
|
const int cta_base_row = blockIdx.x * ROWS_PER_CTA;
|
||||||
|
|
||||||
|
// Now, using the base row per thread block, we compute the base row per warp.
|
||||||
|
const int warp_base_row = cta_base_row + threadIdx.y * ROWS_PER_WARP;
|
||||||
|
|
||||||
|
// The threads in a warp are split into sub-groups that will work on a row.
|
||||||
|
// We compute row offset for each thread sub-group
|
||||||
|
const int thread_row_in_warp = threadIdx.x / THREADS_PER_ROW;
|
||||||
|
const int thread_row = warp_base_row + thread_row_in_warp;
|
||||||
|
|
||||||
|
// Threads with indices out of bounds should early exit here.
|
||||||
|
if (thread_row >= num_rows) {
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
const bool row_is_active = finished ? !finished[thread_row] : true;
|
||||||
|
|
||||||
|
// We finally start setting up the read pointers for each thread. First, each
|
||||||
|
// thread jumps to the start of the row it will read.
|
||||||
|
const T* thread_row_ptr = input + thread_row * ELTS_PER_ROW;
|
||||||
|
|
||||||
|
// Now, we compute the group each thread belong to in order to determine the
|
||||||
|
// first column to start loads.
|
||||||
|
const int thread_group_idx = threadIdx.x % THREADS_PER_ROW;
|
||||||
|
const int first_elt_read_by_thread = thread_group_idx * ELTS_PER_LDG;
|
||||||
|
const T* thread_read_ptr = thread_row_ptr + first_elt_read_by_thread;
|
||||||
|
|
||||||
|
// Determine the pointer type to use to read in the data depending on the
|
||||||
|
// BYTES_PER_LDG template param. In theory, this can support all powers of 2
|
||||||
|
// up to 16. NOTE(woosuk): The original implementation uses CUTLASS aligned
|
||||||
|
// array here. We defined our own aligned array and use it here to avoid the
|
||||||
|
// dependency on CUTLASS.
|
||||||
|
using AccessType = AlignedArray<T, ELTS_PER_LDG>;
|
||||||
|
|
||||||
|
// Finally, we pull in the data from global mem
|
||||||
|
T row_chunk_temp[VPT];
|
||||||
|
AccessType* row_chunk_vec_ptr =
|
||||||
|
reinterpret_cast<AccessType*>(&row_chunk_temp);
|
||||||
|
const AccessType* vec_thread_read_ptr =
|
||||||
|
reinterpret_cast<const AccessType*>(thread_read_ptr);
|
||||||
|
#pragma unroll
|
||||||
|
// Note(Byron): interleaved loads to achieve better memory coalescing
|
||||||
|
// | thread[0] | thread[1] | thread[2] | thread[3] | thread[0] | thread[1] |
|
||||||
|
// thread[2] | thread[3] | ...
|
||||||
|
for (int ii = 0; ii < LDG_PER_THREAD; ++ii) {
|
||||||
|
row_chunk_vec_ptr[ii] = vec_thread_read_ptr[ii * THREADS_PER_ROW];
|
||||||
|
}
|
||||||
|
|
||||||
|
float row_chunk[VPT];
|
||||||
|
#pragma unroll
|
||||||
|
// Note(Byron): upcast logits to float32
|
||||||
|
for (int ii = 0; ii < VPT; ++ii) {
|
||||||
|
row_chunk[ii] = convert_to_float<T>(row_chunk_temp[ii]);
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply tanh softcapping and correction bias
|
||||||
|
if (moe_softcapping != 0.0f || correction_bias != nullptr) {
|
||||||
|
#pragma unroll
|
||||||
|
for (int ii = 0; ii < VPT; ++ii) {
|
||||||
|
float val = row_chunk[ii];
|
||||||
|
|
||||||
|
// Apply tanh softcapping if enabled
|
||||||
|
if (moe_softcapping != 0.0f) {
|
||||||
|
val = tanhf(val / moe_softcapping) * moe_softcapping;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Apply correction bias if provided
|
||||||
|
if (correction_bias != nullptr) {
|
||||||
|
/*
|
||||||
|
LDG is interleaved
|
||||||
|
|thread0 LDG| |thread1 LDG| |thread0 LDG| |thread1 LDG|
|
||||||
|
|--------- group0 --------| |----------group1 --------|
|
||||||
|
^ local2
|
||||||
|
*/
|
||||||
|
const int group_id = ii / ELTS_PER_LDG;
|
||||||
|
const int local_id = ii % ELTS_PER_LDG;
|
||||||
|
const int expert_idx = first_elt_read_by_thread +
|
||||||
|
group_id * THREADS_PER_ROW * ELTS_PER_LDG +
|
||||||
|
local_id;
|
||||||
|
val = val + correction_bias[expert_idx];
|
||||||
|
}
|
||||||
|
|
||||||
|
row_chunk[ii] = val;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// First, we perform a max reduce within the thread. We can do the max in fp16
|
||||||
|
// safely (I think) and just convert to float afterwards for the exp + sum
|
||||||
|
// reduction.
|
||||||
|
float thread_max = row_chunk[0];
|
||||||
|
#pragma unroll
|
||||||
|
for (int ii = 1; ii < VPT; ++ii) {
|
||||||
|
thread_max = max(thread_max, row_chunk[ii]);
|
||||||
|
}
|
||||||
|
|
||||||
|
/*********************************/
|
||||||
|
/********* Softmax Begin *********/
|
||||||
|
/*********************************/
|
||||||
|
|
||||||
|
// Now, we find the max within the thread group and distribute among the
|
||||||
|
// threads. We use a butterfly reduce. lane id: 0-31 within a warp
|
||||||
|
#pragma unroll
|
||||||
|
for (int mask = THREADS_PER_ROW / 2; mask > 0; mask /= 2) {
|
||||||
|
// butterfly reduce with (lane id ^ mask)
|
||||||
|
thread_max = max(thread_max,
|
||||||
|
XLLM_SHFL_XOR_SYNC_WIDTH(
|
||||||
|
0xffffffff, thread_max, mask, THREADS_PER_ROW));
|
||||||
|
}
|
||||||
|
|
||||||
|
// From this point, thread max in all the threads have the max within the row.
|
||||||
|
// Now, we subtract the max from each element in the thread and take the exp.
|
||||||
|
// We also compute the thread local sum.
|
||||||
|
float row_sum = 0;
|
||||||
|
#pragma unroll
|
||||||
|
for (int ii = 0; ii < VPT; ++ii) {
|
||||||
|
row_chunk[ii] = expf(row_chunk[ii] - thread_max);
|
||||||
|
row_sum += row_chunk[ii];
|
||||||
|
}
|
||||||
|
|
||||||
|
// Now, we perform the sum reduce within each thread group. Similar to the max
|
||||||
|
// reduce, we use a bufferfly pattern.
|
||||||
|
#pragma unroll
|
||||||
|
for (int mask = THREADS_PER_ROW / 2; mask > 0; mask /= 2) {
|
||||||
|
row_sum +=
|
||||||
|
XLLM_SHFL_XOR_SYNC_WIDTH(0xffffffff, row_sum, mask, THREADS_PER_ROW);
|
||||||
|
}
|
||||||
|
|
||||||
|
// From this point, all threads have the max and the sum for their rows in the
|
||||||
|
// thread_max and thread_sum variables respectively. Finally, we can scale the
|
||||||
|
// rows for the softmax. Technically, for top-k gating we don't need to
|
||||||
|
// compute the entire softmax row. We can likely look at the maxes and only
|
||||||
|
// compute for the top-k values in the row. However, this kernel will likely
|
||||||
|
// not be a bottle neck and it seems better to closer match torch and find the
|
||||||
|
// argmax after computing the softmax.
|
||||||
|
const float reciprocal_row_sum = 1.f / row_sum;
|
||||||
|
|
||||||
|
#pragma unroll
|
||||||
|
for (int ii = 0; ii < VPT; ++ii) {
|
||||||
|
row_chunk[ii] = row_chunk[ii] * reciprocal_row_sum;
|
||||||
|
}
|
||||||
|
/*******************************/
|
||||||
|
/********* Softmax End *********/
|
||||||
|
/*******************************/
|
||||||
|
|
||||||
|
// Now, softmax_res contains the softmax of the row chunk. Now, I want to find
|
||||||
|
// the topk elements in each row, along with the max index.
|
||||||
|
int start_col = first_elt_read_by_thread;
|
||||||
|
static constexpr int COLS_PER_GROUP_LDG = ELTS_PER_LDG * THREADS_PER_ROW;
|
||||||
|
|
||||||
|
float row_sum_for_renormalize = 0;
|
||||||
|
|
||||||
|
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||||
|
// First, each thread does the local argmax
|
||||||
|
float max_val = row_chunk[0];
|
||||||
|
int expert = start_col;
|
||||||
|
#pragma unroll
|
||||||
|
for (int ldg = 0, col = start_col; ldg < LDG_PER_THREAD;
|
||||||
|
++ldg, col += COLS_PER_GROUP_LDG) {
|
||||||
|
#pragma unroll
|
||||||
|
for (int ii = 0; ii < ELTS_PER_LDG; ++ii) {
|
||||||
|
float val = row_chunk[ldg * ELTS_PER_LDG + ii];
|
||||||
|
|
||||||
|
// No check on the experts here since columns with the smallest index
|
||||||
|
// are processed first and only updated if > (not >=)
|
||||||
|
if (val > max_val) {
|
||||||
|
max_val = val;
|
||||||
|
expert = col + ii;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Now, we perform the argmax reduce. We use the butterfly pattern so threads
|
||||||
|
// reach consensus about the max. This will be useful for K > 1 so that the
|
||||||
|
// threads can agree on "who" had the max value. That thread can then blank out
|
||||||
|
// their max with -inf and the warp can run more iterations...
|
||||||
|
#pragma unroll
|
||||||
|
for (int mask = THREADS_PER_ROW / 2; mask > 0; mask /= 2) {
|
||||||
|
float other_max =
|
||||||
|
XLLM_SHFL_XOR_SYNC_WIDTH(0xffffffff, max_val, mask, THREADS_PER_ROW);
|
||||||
|
int other_expert =
|
||||||
|
XLLM_SHFL_XOR_SYNC_WIDTH(0xffffffff, expert, mask, THREADS_PER_ROW);
|
||||||
|
|
||||||
|
// We want lower indices to "win" in every thread so we break ties this
|
||||||
|
// way
|
||||||
|
if (other_max > max_val ||
|
||||||
|
(other_max == max_val && other_expert < expert)) {
|
||||||
|
max_val = other_max;
|
||||||
|
expert = other_expert;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Write the max for this k iteration to global memory.
|
||||||
|
if (thread_group_idx == 0) {
|
||||||
|
// Add a guard to ignore experts not included by this node
|
||||||
|
const bool node_uses_expert =
|
||||||
|
expert >= start_expert && expert < end_expert;
|
||||||
|
const bool should_process_row = row_is_active && node_uses_expert;
|
||||||
|
|
||||||
|
// The lead thread from each sub-group will write out the final results to
|
||||||
|
// global memory. (This will be a single) thread per row of the
|
||||||
|
// input/output matrices.
|
||||||
|
const int idx = k * thread_row + k_idx;
|
||||||
|
output[idx] = max_val;
|
||||||
|
indices[idx] = should_process_row ? (expert - start_expert) : NUM_EXPERTS;
|
||||||
|
row_sum_for_renormalize += max_val;
|
||||||
|
}
|
||||||
|
|
||||||
|
// Finally, we clear the value in the thread with the current max if there
|
||||||
|
// is another iteration to run.
|
||||||
|
if (k_idx + 1 < k) {
|
||||||
|
const int ldg_group_for_expert = expert / COLS_PER_GROUP_LDG;
|
||||||
|
const int thread_to_clear_in_group =
|
||||||
|
(expert / ELTS_PER_LDG) % THREADS_PER_ROW;
|
||||||
|
|
||||||
|
// Only the thread in the group which produced the max will reset the
|
||||||
|
// "winning" value to -inf.
|
||||||
|
if (thread_group_idx == thread_to_clear_in_group) {
|
||||||
|
const int offset_for_expert = expert % ELTS_PER_LDG;
|
||||||
|
// Safe to set to any negative value since row_chunk values must be
|
||||||
|
// between 0 and 1.
|
||||||
|
row_chunk[ldg_group_for_expert * ELTS_PER_LDG + offset_for_expert] =
|
||||||
|
-10000.f;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Fuse renormalization of topk_weights into this kernel
|
||||||
|
if (renormalize && thread_group_idx == 0) {
|
||||||
|
float row_sum_for_renormalize_inv = 1.f / row_sum_for_renormalize;
|
||||||
|
#pragma unroll
|
||||||
|
for (int k_idx = 0; k_idx < k; ++k_idx) {
|
||||||
|
const int idx = k * thread_row + k_idx;
|
||||||
|
output[idx] = output[idx] * row_sum_for_renormalize_inv;
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
template <typename T, int EXPERTS, int WARPS_PER_TB>
|
||||||
|
void topk_gating_softmax_launcher_helper(const T* input,
|
||||||
|
const bool* finished,
|
||||||
|
float* output,
|
||||||
|
int* indices,
|
||||||
|
const int num_rows,
|
||||||
|
const int k,
|
||||||
|
const int start_expert,
|
||||||
|
const int end_expert,
|
||||||
|
const bool renormalize,
|
||||||
|
const float moe_softcapping,
|
||||||
|
const float* correction_bias,
|
||||||
|
cudaStream_t stream) {
|
||||||
|
static constexpr std::size_t MAX_BYTES_PER_LDG = 16;
|
||||||
|
|
||||||
|
static constexpr int BYTES_PER_LDG =
|
||||||
|
MIN(MAX_BYTES_PER_LDG, sizeof(T) * EXPERTS);
|
||||||
|
using Constants = TopkConstants<T, EXPERTS, BYTES_PER_LDG>;
|
||||||
|
static constexpr int VPT = Constants::VPT;
|
||||||
|
static constexpr int ROWS_PER_WARP = Constants::ROWS_PER_WARP;
|
||||||
|
const int num_warps = (num_rows + ROWS_PER_WARP - 1) / ROWS_PER_WARP;
|
||||||
|
const int num_blocks = (num_warps + WARPS_PER_TB - 1) / WARPS_PER_TB;
|
||||||
|
|
||||||
|
dim3 block_dim(WARP_SIZE, WARPS_PER_TB);
|
||||||
|
topk_gating_softmax<T, VPT, EXPERTS, WARPS_PER_TB, BYTES_PER_LDG>
|
||||||
|
<<<num_blocks, block_dim, 0, stream>>>(input,
|
||||||
|
finished,
|
||||||
|
output,
|
||||||
|
num_rows,
|
||||||
|
indices,
|
||||||
|
k,
|
||||||
|
start_expert,
|
||||||
|
end_expert,
|
||||||
|
renormalize,
|
||||||
|
moe_softcapping,
|
||||||
|
correction_bias);
|
||||||
|
}
|
||||||
|
|
||||||
|
#define LAUNCH_SOFTMAX(TYPE, NUM_EXPERTS, WARPS_PER_TB) \
|
||||||
|
topk_gating_softmax_launcher_helper<TYPE, NUM_EXPERTS, WARPS_PER_TB>( \
|
||||||
|
gating_output, \
|
||||||
|
nullptr, \
|
||||||
|
topk_weights, \
|
||||||
|
topk_indices, \
|
||||||
|
num_tokens, \
|
||||||
|
topk, \
|
||||||
|
0, \
|
||||||
|
num_experts, \
|
||||||
|
renormalize, \
|
||||||
|
moe_softcapping, \
|
||||||
|
correction_bias, \
|
||||||
|
stream);
|
||||||
|
|
||||||
|
template <typename T>
|
||||||
|
void topk_gating_softmax_kernel_launcher(const T* gating_output,
|
||||||
|
float* topk_weights,
|
||||||
|
int* topk_indices,
|
||||||
|
float* softmax_workspace,
|
||||||
|
const int num_tokens,
|
||||||
|
const int num_experts,
|
||||||
|
const int topk,
|
||||||
|
const bool renormalize,
|
||||||
|
const float moe_softcapping,
|
||||||
|
const float* correction_bias,
|
||||||
|
cudaStream_t stream) {
|
||||||
|
static constexpr int WARPS_PER_TB = 4;
|
||||||
|
switch (num_experts) {
|
||||||
|
case 1:
|
||||||
|
LAUNCH_SOFTMAX(T, 1, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 2:
|
||||||
|
LAUNCH_SOFTMAX(T, 2, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 4:
|
||||||
|
LAUNCH_SOFTMAX(T, 4, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 8:
|
||||||
|
LAUNCH_SOFTMAX(T, 8, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 16:
|
||||||
|
LAUNCH_SOFTMAX(T, 16, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 32:
|
||||||
|
LAUNCH_SOFTMAX(T, 32, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 64:
|
||||||
|
LAUNCH_SOFTMAX(T, 64, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 128:
|
||||||
|
LAUNCH_SOFTMAX(T, 128, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
case 256:
|
||||||
|
LAUNCH_SOFTMAX(T, 256, WARPS_PER_TB);
|
||||||
|
break;
|
||||||
|
default: {
|
||||||
|
TORCH_CHECK(softmax_workspace != nullptr, "softmax_workspace must be provided for num_experts that are not a power of 2.");
|
||||||
|
static constexpr int TPB = 256;
|
||||||
|
moe_softmax<T, TPB><<<num_tokens, TPB, 0, stream>>>(gating_output,
|
||||||
|
nullptr,
|
||||||
|
softmax_workspace,
|
||||||
|
num_experts,
|
||||||
|
moe_softcapping,
|
||||||
|
correction_bias);
|
||||||
|
if (topk == 1) {
|
||||||
|
// Note: As an optimization for better performance,
|
||||||
|
// the softmax_workspace is overwritten in-place by both moeTopK and
|
||||||
|
// moe_topk_fast.
|
||||||
|
moe_topK<TPB><<<num_tokens, TPB, 0, stream>>>(softmax_workspace,
|
||||||
|
nullptr,
|
||||||
|
topk_weights,
|
||||||
|
topk_indices,
|
||||||
|
num_experts,
|
||||||
|
topk,
|
||||||
|
0,
|
||||||
|
num_experts,
|
||||||
|
renormalize);
|
||||||
|
} else {
|
||||||
|
moe_topk_fast<TPB><<<num_tokens, TPB, 0, stream>>>(softmax_workspace,
|
||||||
|
nullptr,
|
||||||
|
topk_weights,
|
||||||
|
topk_indices,
|
||||||
|
num_experts,
|
||||||
|
topk,
|
||||||
|
0,
|
||||||
|
num_experts,
|
||||||
|
renormalize);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} // namespace
|
||||||
|
|
||||||
|
namespace xllm::kernel::cuda {
|
||||||
|
void topk_softmax(torch::Tensor& topk_weights, // [num_tokens, topk]
|
||||||
|
torch::Tensor& topk_indices, // [num_tokens, topk]
|
||||||
|
torch::Tensor& gating_output, // [num_tokens, num_experts]
|
||||||
|
const bool renormalize,
|
||||||
|
const double moe_softcapping,
|
||||||
|
const std::optional<torch::Tensor>& correction_bias) {
|
||||||
|
// Check data type
|
||||||
|
TORCH_CHECK(gating_output.scalar_type() == at::ScalarType::Float ||
|
||||||
|
gating_output.scalar_type() == at::ScalarType::Half ||
|
||||||
|
gating_output.scalar_type() == at::ScalarType::BFloat16,
|
||||||
|
"gating_output must be float32, float16, or bfloat16");
|
||||||
|
|
||||||
|
// Check dimensions
|
||||||
|
TORCH_CHECK(gating_output.dim() == 2, "gating_output must be 2D tensor [num_tokens, num_experts]");
|
||||||
|
TORCH_CHECK(topk_weights.dim() == 2, "topk_weights must be 2D tensor [num_tokens, topk]");
|
||||||
|
TORCH_CHECK(topk_indices.dim() == 2, "topk_indices must be 2D tensor [num_tokens, topk]");
|
||||||
|
|
||||||
|
// Check shapes
|
||||||
|
TORCH_CHECK(gating_output.size(0) == topk_weights.size(0), "First dimension of topk_weights must match num_tokens in gating_output First dimension of topk_indices must match num_tokens in gating_output");
|
||||||
|
|
||||||
|
TORCH_CHECK(topk_weights.size(-1) == topk_indices.size(-1), "Second dimension of topk_indices must match topk in topk_weights topk must be less than or equal to num_experts");
|
||||||
|
|
||||||
|
const int num_experts = static_cast<int>(gating_output.size(-1));
|
||||||
|
const int num_tokens = static_cast<int>(gating_output.size(0));
|
||||||
|
const int topk = static_cast<int>(topk_weights.size(-1));
|
||||||
|
|
||||||
|
const bool is_pow_2 =
|
||||||
|
(num_experts != 0) && ((num_experts & (num_experts - 1)) == 0);
|
||||||
|
const bool needs_workspace = !is_pow_2 || num_experts > 256;
|
||||||
|
const int64_t workspace_size = needs_workspace ? num_tokens * num_experts : 0;
|
||||||
|
|
||||||
|
const at::cuda::OptionalCUDAGuard device_guard(device_of(gating_output));
|
||||||
|
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
|
||||||
|
torch::Tensor softmax_workspace = torch::empty(
|
||||||
|
{workspace_size}, gating_output.options().dtype(at::ScalarType::Float));
|
||||||
|
|
||||||
|
const at::ScalarType dtype = gating_output.scalar_type();
|
||||||
|
|
||||||
|
// Validate correction_bias if provided - must always be float32
|
||||||
|
const float* bias_ptr = nullptr;
|
||||||
|
if (correction_bias.has_value()) {
|
||||||
|
const torch::Tensor& bias_tensor = correction_bias.value();
|
||||||
|
TORCH_CHECK(bias_tensor.dim() == 1, "correction_bias must be 1D tensor [num_experts]");
|
||||||
|
TORCH_CHECK(bias_tensor.size(0) == num_experts, "correction_bias size must match num_experts");
|
||||||
|
TORCH_CHECK(bias_tensor.scalar_type() == at::ScalarType::Float, "correction_bias must be float32");
|
||||||
|
bias_ptr = bias_tensor.data_ptr<float>();
|
||||||
|
}
|
||||||
|
|
||||||
|
// Cast moe_softcapping from double to float for CUDA kernels
|
||||||
|
const float moe_softcapping_f = static_cast<float>(moe_softcapping);
|
||||||
|
|
||||||
|
if (dtype == at::ScalarType::Float) {
|
||||||
|
topk_gating_softmax_kernel_launcher<float>(
|
||||||
|
gating_output.data_ptr<float>(),
|
||||||
|
topk_weights.data_ptr<float>(),
|
||||||
|
topk_indices.data_ptr<int>(),
|
||||||
|
softmax_workspace.data_ptr<float>(),
|
||||||
|
num_tokens,
|
||||||
|
num_experts,
|
||||||
|
topk,
|
||||||
|
renormalize,
|
||||||
|
moe_softcapping_f,
|
||||||
|
bias_ptr,
|
||||||
|
stream);
|
||||||
|
} else if (dtype == at::ScalarType::Half) {
|
||||||
|
topk_gating_softmax_kernel_launcher<__half>(
|
||||||
|
reinterpret_cast<const __half*>(gating_output.data_ptr<at::Half>()),
|
||||||
|
topk_weights.data_ptr<float>(),
|
||||||
|
topk_indices.data_ptr<int>(),
|
||||||
|
softmax_workspace.data_ptr<float>(),
|
||||||
|
num_tokens,
|
||||||
|
num_experts,
|
||||||
|
topk,
|
||||||
|
renormalize,
|
||||||
|
moe_softcapping_f,
|
||||||
|
bias_ptr,
|
||||||
|
stream);
|
||||||
|
} else if (dtype == at::ScalarType::BFloat16) {
|
||||||
|
topk_gating_softmax_kernel_launcher<__nv_bfloat16>(
|
||||||
|
reinterpret_cast<const __nv_bfloat16*>(
|
||||||
|
gating_output.data_ptr<at::BFloat16>()),
|
||||||
|
topk_weights.data_ptr<float>(),
|
||||||
|
topk_indices.data_ptr<int>(),
|
||||||
|
softmax_workspace.data_ptr<float>(),
|
||||||
|
num_tokens,
|
||||||
|
num_experts,
|
||||||
|
topk,
|
||||||
|
renormalize,
|
||||||
|
moe_softcapping_f,
|
||||||
|
bias_ptr,
|
||||||
|
stream);
|
||||||
|
} else {
|
||||||
|
TORCH_CHECK(false, "Unsupported gating_output dtype");
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} // namespace xllm::kernel::cuda
|
||||||
BIN
qwen3_6_scripts/prebuilt/corex-3.2.3-ivcore10/corex_moe_topk_softmax.so
Executable file
BIN
qwen3_6_scripts/prebuilt/corex-3.2.3-ivcore10/corex_moe_topk_softmax.so
Executable file
Binary file not shown.
@@ -180,6 +180,8 @@ class ChatCompletionRequest(OpenAIBaseModel):
|
|||||||
logprobs: Optional[bool] = False
|
logprobs: Optional[bool] = False
|
||||||
top_logprobs: Optional[int] = 0
|
top_logprobs: Optional[int] = 0
|
||||||
max_tokens: Optional[int] = None
|
max_tokens: Optional[int] = None
|
||||||
|
# OpenAI newer API field — treat as alias for max_tokens
|
||||||
|
max_completion_tokens: Optional[int] = None
|
||||||
n: Optional[int] = 1
|
n: Optional[int] = 1
|
||||||
presence_penalty: Optional[float] = 0.0
|
presence_penalty: Optional[float] = 0.0
|
||||||
response_format: Optional[ResponseFormat] = None
|
response_format: Optional[ResponseFormat] = None
|
||||||
@@ -193,6 +195,7 @@ class ChatCompletionRequest(OpenAIBaseModel):
|
|||||||
tool_choice: Optional[Union[Literal["none"], Literal["auto"],
|
tool_choice: Optional[Union[Literal["none"], Literal["auto"],
|
||||||
ChatCompletionNamedToolChoiceParam]] = "none"
|
ChatCompletionNamedToolChoiceParam]] = "none"
|
||||||
thinking: Optional[Union[bool, str, Dict[str, Any]]] = None
|
thinking: Optional[Union[bool, str, Dict[str, Any]]] = None
|
||||||
|
reasoning_effort: Optional[str] = None
|
||||||
|
|
||||||
# NOTE this will be ignored by VLLM -- the model determines the behavior
|
# NOTE this will be ignored by VLLM -- the model determines the behavior
|
||||||
parallel_tool_calls: Optional[bool] = False
|
parallel_tool_calls: Optional[bool] = False
|
||||||
|
|||||||
@@ -138,6 +138,11 @@ try:
|
|||||||
except ImportError:
|
except ImportError:
|
||||||
_corex_moe_direct_routed = None
|
_corex_moe_direct_routed = None
|
||||||
|
|
||||||
|
try:
|
||||||
|
from vllm import corex_moe_topk_softmax as _corex_moe_topk_softmax
|
||||||
|
except ImportError:
|
||||||
|
_corex_moe_topk_softmax = None
|
||||||
|
|
||||||
from vllm.model_executor.models.interfaces import (HasInnerState, SupportsLoRA,
|
from vllm.model_executor.models.interfaces import (HasInnerState, SupportsLoRA,
|
||||||
SupportsMultiModal)
|
SupportsMultiModal)
|
||||||
|
|
||||||
@@ -179,6 +184,9 @@ _USE_COREX_MOE_WEIGHT_GATHER = (
|
|||||||
_USE_COREX_MOE_DIRECT_ROUTED = (
|
_USE_COREX_MOE_DIRECT_ROUTED = (
|
||||||
_corex_moe_direct_routed is not None
|
_corex_moe_direct_routed is not None
|
||||||
and env_bool("BI100_MOE_COREX_DIRECT_ROUTED", False))
|
and env_bool("BI100_MOE_COREX_DIRECT_ROUTED", False))
|
||||||
|
_USE_COREX_MOE_TOPK_SOFTMAX = (
|
||||||
|
_corex_moe_topk_softmax is not None
|
||||||
|
and env_bool("BI100_MOE_COREX_TOPK_SOFTMAX", True))
|
||||||
_USE_FUSED_MOE_ACTIVATION = env_bool("BI100_MOE_FUSED_ACTIVATION", True)
|
_USE_FUSED_MOE_ACTIVATION = env_bool("BI100_MOE_FUSED_ACTIVATION", True)
|
||||||
|
|
||||||
|
|
||||||
@@ -1599,13 +1607,18 @@ class Qwen3_5MoeSparseBlock(nn.Module):
|
|||||||
Output is partial (pre-all-reduce), same contract as FusedMoE
|
Output is partial (pre-all-reduce), same contract as FusedMoE
|
||||||
with reduce_results=False.
|
with reduce_results=False.
|
||||||
"""
|
"""
|
||||||
# Softmax is monotonic, so selecting logits first is equivalent to
|
# Fused topk+softmax: single CUB kernel vs 2 PyTorch ops.
|
||||||
# full-expert softmax -> top-k -> renormalise while normalising only K
|
# Source: xllm/core/kernels/cuda/moe/moe_topk_softmax_kernels.cuh
|
||||||
# values. This saves one 256-wide softmax in the decode hot path.
|
if _USE_COREX_MOE_TOPK_SOFTMAX:
|
||||||
topk_logits, topk_ids = torch.topk(
|
topk_weights, topk_ids = _corex_moe_topk_softmax.moe_topk_softmax(
|
||||||
router_logits.float(), self.top_k, dim=-1) # (T, top_k)
|
router_logits, self.top_k, True)
|
||||||
topk_weights = torch.softmax(topk_logits, dim=-1)
|
topk_ids = topk_ids.to(torch.int64)
|
||||||
topk_weights = topk_weights.to(hidden_states.dtype)
|
topk_weights = topk_weights.to(hidden_states.dtype)
|
||||||
|
else:
|
||||||
|
topk_logits, topk_ids = torch.topk(
|
||||||
|
router_logits.float(), self.top_k, dim=-1) # (T, top_k)
|
||||||
|
topk_weights = torch.softmax(topk_logits, dim=-1)
|
||||||
|
topk_weights = topk_weights.to(hidden_states.dtype)
|
||||||
|
|
||||||
w13 = self.experts.w13_weight # (E, 2*I, H)
|
w13 = self.experts.w13_weight # (E, 2*I, H)
|
||||||
w2 = self.experts.w2_weight # (E, H, I)
|
w2 = self.experts.w2_weight # (E, H, I)
|
||||||
|
|||||||
@@ -394,6 +394,9 @@ class OpenAIServingChat(OpenAIServing):
|
|||||||
assert prompt_inputs is not None
|
assert prompt_inputs is not None
|
||||||
|
|
||||||
sampling_params: Union[SamplingParams, BeamSearchParams]
|
sampling_params: Union[SamplingParams, BeamSearchParams]
|
||||||
|
# OpenAI API: max_completion_tokens takes precedence over max_tokens
|
||||||
|
if request.max_completion_tokens is not None and request.max_tokens is None:
|
||||||
|
request.max_tokens = request.max_completion_tokens
|
||||||
default_max_tokens = self.max_model_len - len(
|
default_max_tokens = self.max_model_len - len(
|
||||||
prompt_inputs["prompt_token_ids"])
|
prompt_inputs["prompt_token_ids"])
|
||||||
if request.use_beam_search:
|
if request.use_beam_search:
|
||||||
|
|||||||
233
vllm/corex_moe.py
Normal file
233
vllm/corex_moe.py
Normal file
@@ -0,0 +1,233 @@
|
|||||||
|
"""
|
||||||
|
corex_moe.py — Fused MoE dispatch for BI-V100 via ix_moe_bridge.so
|
||||||
|
|
||||||
|
Sub168 log reference:
|
||||||
|
corex_moe.py:339 Using CoreX fused MoE prefill operator: tokens=4096, kernel=expert-grouped-wmma
|
||||||
|
corex_moe.py:249 Using CoreX fused MoE decode operator
|
||||||
|
|
||||||
|
Call chain:
|
||||||
|
qwen3_5.py → FusedMoE.forward() → corex_moe.forward()
|
||||||
|
→ ix_moe_bridge.topk_softmax() (Step 1: routing)
|
||||||
|
→ ix_moe_bridge.moe_gen_idx() (Step 2: index generation)
|
||||||
|
→ ix_moe_bridge.moe_expand_input() (Step 3: expand)
|
||||||
|
→ ix_moe_bridge.moe_group_gemm() (Step 4: w13 gate+up GEMM)
|
||||||
|
→ ix_moe_bridge.silu_and_mul() (Step 5: activation)
|
||||||
|
→ ix_moe_bridge.moe_group_gemm() (Step 6: w2 down GEMM)
|
||||||
|
→ ix_moe_bridge.moe_combine_result() (Step 7: weighted sum)
|
||||||
|
|
||||||
|
Source: upstream_ref/xllm/xllm/core/kernels/ilu/fused_moe.cpp
|
||||||
|
upstream_ref/xllm/xllm/core/kernels/ilu/ixformer.h
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import glob
|
||||||
|
import torch
|
||||||
|
from typing import Optional, Tuple
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
# ============================================================================
|
||||||
|
# Load ix_moe_bridge.so — compiled by precompile_ix_bridge.py in Docker
|
||||||
|
# ============================================================================
|
||||||
|
_bridge = None
|
||||||
|
_bridge_load_attempted = False
|
||||||
|
|
||||||
|
|
||||||
|
def _load_bridge():
|
||||||
|
"""Try to load ix_moe_bridge.so from known paths."""
|
||||||
|
global _bridge, _bridge_load_attempted
|
||||||
|
if _bridge_load_attempted:
|
||||||
|
return _bridge
|
||||||
|
_bridge_load_attempted = True
|
||||||
|
|
||||||
|
search_paths = [
|
||||||
|
"/usr/local/corex/lib/python3/dist-packages/ex_engine/build",
|
||||||
|
"/usr/local/corex/lib/python3/dist-packages/ex_engine",
|
||||||
|
"/usr/local/corex/lib/python3/dist-packages",
|
||||||
|
"/workspace/ex_engine/build",
|
||||||
|
"/workspace/ex_engine",
|
||||||
|
]
|
||||||
|
|
||||||
|
for d in search_paths:
|
||||||
|
for so in glob.glob(os.path.join(d, "ix_moe_bridge*.so")):
|
||||||
|
try:
|
||||||
|
import importlib.util
|
||||||
|
spec = importlib.util.spec_from_file_location("ix_moe_bridge", so)
|
||||||
|
mod = importlib.util.module_from_spec(spec)
|
||||||
|
spec.loader.exec_module(mod)
|
||||||
|
_bridge = mod
|
||||||
|
logger.info("Loaded ix_moe_bridge from %s", so)
|
||||||
|
return _bridge
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug("Failed loading %s: %s", so, e)
|
||||||
|
|
||||||
|
# Fallback: try torch.ops (if registered via JIT during build)
|
||||||
|
try:
|
||||||
|
import torch.utils.cpp_extension
|
||||||
|
_bridge = torch.utils.cpp_extension.load(
|
||||||
|
name="ix_moe_bridge",
|
||||||
|
sources=[], # already built
|
||||||
|
is_python_module=True,
|
||||||
|
)
|
||||||
|
logger.info("Loaded ix_moe_bridge via torch extension cache")
|
||||||
|
return _bridge
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
logger.warning("ix_moe_bridge.so not found — MoE will use PyTorch fallback (SLOW)")
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
class CoreXMoE:
|
||||||
|
"""
|
||||||
|
Fused MoE operator matching qwen3_5.py FusedMoE call convention.
|
||||||
|
|
||||||
|
Interface:
|
||||||
|
forward(hidden_states, router_logits, w13, w2, topk, renormalize,
|
||||||
|
num_expert_groups=0, topk_group=0, n_shared_experts=0,
|
||||||
|
shared_expert_gate=None, shared_w13=None, shared_w2=None)
|
||||||
|
→ (output, shared_expert_output_or_None)
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, num_experts: int = 64, topk: int = 8):
|
||||||
|
self.num_experts = num_experts
|
||||||
|
self.topk = topk
|
||||||
|
self._bridge = _load_bridge()
|
||||||
|
self._prefill_logged = False
|
||||||
|
self._decode_logged = False
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
hidden_states: torch.Tensor, # (num_tokens, hidden_size)
|
||||||
|
router_logits: torch.Tensor, # (num_tokens, num_experts)
|
||||||
|
w13: torch.Tensor, # (num_local_experts, 2*intermediate, hidden)
|
||||||
|
w2: torch.Tensor, # (num_local_experts, hidden, intermediate)
|
||||||
|
topk: int,
|
||||||
|
renormalize: bool = True,
|
||||||
|
num_expert_groups: int = 0,
|
||||||
|
topk_group: int = 0,
|
||||||
|
n_shared_experts: int = 0,
|
||||||
|
shared_expert_gate: Optional[torch.Tensor] = None,
|
||||||
|
shared_w13: Optional[torch.Tensor] = None,
|
||||||
|
shared_w2: Optional[torch.Tensor] = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Full fused MoE forward via ixformer C++ bridge."""
|
||||||
|
|
||||||
|
num_tokens = hidden_states.size(0)
|
||||||
|
hidden_size = hidden_states.size(1)
|
||||||
|
num_local_experts = w13.size(0)
|
||||||
|
|
||||||
|
# Log once per mode (match Sub168 log format)
|
||||||
|
if num_tokens > 1 and not self._prefill_logged:
|
||||||
|
logger.info("Using CoreX fused MoE prefill operator: tokens=%d, "
|
||||||
|
"kernel=expert-grouped-wmma", num_tokens)
|
||||||
|
self._prefill_logged = True
|
||||||
|
elif num_tokens == 1 and not self._decode_logged:
|
||||||
|
logger.info("Using CoreX fused MoE decode operator")
|
||||||
|
self._decode_logged = True
|
||||||
|
|
||||||
|
if self._bridge is not None:
|
||||||
|
return self._forward_bridge(
|
||||||
|
hidden_states, router_logits, w13, w2, topk,
|
||||||
|
renormalize, num_local_experts, hidden_size)
|
||||||
|
else:
|
||||||
|
return self._forward_pytorch(
|
||||||
|
hidden_states, router_logits, w13, w2, topk,
|
||||||
|
renormalize, num_local_experts, hidden_size)
|
||||||
|
|
||||||
|
def _forward_bridge(
|
||||||
|
self, hidden_states, router_logits, w13, w2,
|
||||||
|
topk, renormalize, num_local_experts, hidden_size
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""7-step fused MoE via ix_moe_bridge.so → ixformer::infer."""
|
||||||
|
bridge = self._bridge
|
||||||
|
num_tokens = hidden_states.size(0)
|
||||||
|
num_experts = router_logits.size(1)
|
||||||
|
|
||||||
|
# Step 1: topk_softmax
|
||||||
|
gating = router_logits.to(torch.float32)
|
||||||
|
topk_weights = torch.empty(
|
||||||
|
(num_tokens, topk), dtype=torch.float32, device=hidden_states.device)
|
||||||
|
topk_ids = torch.empty(
|
||||||
|
(num_tokens, topk), dtype=torch.int32, device=hidden_states.device)
|
||||||
|
token_expert_indices = torch.empty(
|
||||||
|
(num_tokens, topk), dtype=torch.int32, device=hidden_states.device)
|
||||||
|
|
||||||
|
bridge.topk_softmax(topk_weights, topk_ids, token_expert_indices, gating)
|
||||||
|
|
||||||
|
if renormalize:
|
||||||
|
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||||
|
|
||||||
|
# Step 2: generate index
|
||||||
|
idx_result = bridge.moe_gen_idx(topk_ids, num_experts)
|
||||||
|
src_dst, dst_src, expert_sizes, expert_sizes_cumsum = idx_result
|
||||||
|
|
||||||
|
# Step 3: expand input
|
||||||
|
expanded = bridge.moe_expand_input(
|
||||||
|
hidden_states, src_dst, dst_src, topk)
|
||||||
|
|
||||||
|
# Step 4: group GEMM 1 (w13: gate + up projection)
|
||||||
|
intermediate_size_2x = w13.size(1)
|
||||||
|
gemm1_out = expanded.new_empty((expanded.size(0), intermediate_size_2x))
|
||||||
|
expert_sizes_cpu = expert_sizes.cpu()
|
||||||
|
bridge.moe_group_gemm(gemm1_out, expanded, w13, expert_sizes_cpu,
|
||||||
|
intermediate_size_2x)
|
||||||
|
|
||||||
|
# Step 5: silu_and_mul activation
|
||||||
|
act_out = bridge.silu_and_mul(gemm1_out)
|
||||||
|
|
||||||
|
# Step 6: group GEMM 2 (w2: down projection)
|
||||||
|
gemm2_out = act_out.new_empty((act_out.size(0), hidden_size))
|
||||||
|
bridge.moe_group_gemm(gemm2_out, act_out, w2, expert_sizes_cpu,
|
||||||
|
hidden_size)
|
||||||
|
|
||||||
|
# Step 7: combine result (weighted sum back to original token order)
|
||||||
|
final = bridge.moe_combine_result(gemm2_out, topk_weights)
|
||||||
|
|
||||||
|
return final
|
||||||
|
|
||||||
|
def _forward_pytorch(
|
||||||
|
self, hidden_states, router_logits, w13, w2,
|
||||||
|
topk, renormalize, num_local_experts, hidden_size
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Pure PyTorch fallback — SLOW but correct."""
|
||||||
|
num_tokens = hidden_states.size(0)
|
||||||
|
|
||||||
|
# Softmax routing
|
||||||
|
scores = torch.softmax(router_logits.float(), dim=-1)
|
||||||
|
topk_weights, topk_ids = torch.topk(scores, topk, dim=-1)
|
||||||
|
if renormalize:
|
||||||
|
topk_weights = topk_weights / topk_weights.sum(dim=-1, keepdim=True)
|
||||||
|
topk_weights = topk_weights.to(hidden_states.dtype)
|
||||||
|
|
||||||
|
# Expert loop
|
||||||
|
final = torch.zeros(
|
||||||
|
(num_tokens, hidden_size),
|
||||||
|
dtype=hidden_states.dtype, device=hidden_states.device)
|
||||||
|
|
||||||
|
for i in range(num_local_experts):
|
||||||
|
mask = (topk_ids == i).any(dim=-1)
|
||||||
|
if not mask.any():
|
||||||
|
continue
|
||||||
|
idx = mask.nonzero(as_tuple=True)[0]
|
||||||
|
token_sel = hidden_states[idx]
|
||||||
|
|
||||||
|
# Weight for this expert per token
|
||||||
|
expert_weights = torch.zeros(
|
||||||
|
idx.size(0), dtype=topk_weights.dtype, device=hidden_states.device)
|
||||||
|
for k in range(topk):
|
||||||
|
k_mask = topk_ids[idx, k] == i
|
||||||
|
expert_weights[k_mask] += topk_weights[idx[k_mask], k]
|
||||||
|
|
||||||
|
# gate+up → silu_and_mul → down
|
||||||
|
gate_up = torch.mm(token_sel, w13[i].t())
|
||||||
|
half_dim = gate_up.size(-1) // 2
|
||||||
|
gate = gate_up[:, :half_dim]
|
||||||
|
up = gate_up[:, half_dim:]
|
||||||
|
activated = torch.nn.functional.silu(gate) * up
|
||||||
|
down = torch.mm(activated, w2[i].t())
|
||||||
|
|
||||||
|
final[idx] += down * expert_weights.unsqueeze(-1)
|
||||||
|
|
||||||
|
return final
|
||||||
178
vllm/corex_so_loader.py
Normal file
178
vllm/corex_so_loader.py
Normal file
@@ -0,0 +1,178 @@
|
|||||||
|
"""corex_so_loader.py — Unified loader for all 12 prebuilt CoreX .so modules.
|
||||||
|
|
||||||
|
CCCL pattern: device_reduce policy_selector — enumerate available kernels at
|
||||||
|
init, expose a stable Python API, fall back gracefully when .so unavailable.
|
||||||
|
|
||||||
|
The 12 prebuilt .so files expose these operator families:
|
||||||
|
|
||||||
|
GDN decode pipeline (5 .so):
|
||||||
|
corex_gdn_causal_conv → .causal_conv_update(conv_state, mixed_qkv, weight)
|
||||||
|
corex_gdn_packed_decode → .packed_decode(temporal_state, packed_qkv, b, a, A_log, dt_bias)
|
||||||
|
corex_gdn_beta_decay → .beta_decay(b, a, A_log, dt_bias)
|
||||||
|
corex_gdn_qk_map → .qk_map(q, k, num_v_heads)
|
||||||
|
corex_gdn_gated_norm → .apply_inverse(x, z)
|
||||||
|
|
||||||
|
Attention pipeline (3 .so):
|
||||||
|
corex_attn_head_rms_norm → .prepare(x, eps) + .apply_inverse(x, z)
|
||||||
|
corex_paged_kv_gather → .gather(key_cache, val_cache, block_tables, context_lens)
|
||||||
|
corex_fused_paged_prefill → .forward(q, k_cache, v_cache, ...)
|
||||||
|
|
||||||
|
KV cache transfer (1 .so):
|
||||||
|
corex_block_major_kv_transfer → .transfer(src, dst, mapping)
|
||||||
|
|
||||||
|
MoE pipeline (3 .so):
|
||||||
|
corex_moe_direct_routed → .w13(hidden, w13, expert_ids)
|
||||||
|
+ .w2_reduce(act, w2, expert_ids, weights)
|
||||||
|
corex_moe_weight_gather → .gather(w13, w2, expert_ids)
|
||||||
|
corex_moe_exact_reduce → .serial_float(expert_out, weights)
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
from ex_engine.python.corex_so_loader import corex
|
||||||
|
if corex.gdn_causal_conv is not None:
|
||||||
|
out = corex.gdn_causal_conv.causal_conv_update(...)
|
||||||
|
|
||||||
|
# Or import from vllm install root (patch_ops.sh deploys there):
|
||||||
|
from corex_so_loader import corex
|
||||||
|
"""
|
||||||
|
|
||||||
|
import importlib.util
|
||||||
|
import logging
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
logger = logging.getLogger("corex_so_loader")
|
||||||
|
|
||||||
|
# All 12 .so modules in load order
|
||||||
|
_SO_MANIFEST = [
|
||||||
|
"corex_gdn_causal_conv",
|
||||||
|
"corex_gdn_packed_decode",
|
||||||
|
"corex_gdn_beta_decay",
|
||||||
|
"corex_gdn_qk_map",
|
||||||
|
"corex_gdn_gated_norm",
|
||||||
|
"corex_attn_head_rms_norm",
|
||||||
|
"corex_paged_kv_gather",
|
||||||
|
"corex_fused_paged_prefill",
|
||||||
|
"corex_block_major_kv_transfer",
|
||||||
|
"corex_moe_direct_routed",
|
||||||
|
"corex_moe_weight_gather",
|
||||||
|
"corex_moe_exact_reduce",
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def _find_so_dir() -> Optional[str]:
|
||||||
|
"""Find the directory containing prebuilt CoreX .so files.
|
||||||
|
|
||||||
|
Search order:
|
||||||
|
1. COREX_SO_DIR env var
|
||||||
|
2. vllm install roots (where patch_ops.sh installs them)
|
||||||
|
3. Bundled prebuilt directory (repo-relative)
|
||||||
|
4. /usr/local/corex/lib64/
|
||||||
|
"""
|
||||||
|
candidates = []
|
||||||
|
|
||||||
|
env = os.getenv("COREX_SO_DIR")
|
||||||
|
if env:
|
||||||
|
candidates.append(env)
|
||||||
|
|
||||||
|
# vllm install roots (patch_ops.sh copies .so here)
|
||||||
|
for p in sys.path:
|
||||||
|
if "vllm" in p or "dist-packages" in p:
|
||||||
|
candidates.append(p)
|
||||||
|
# Also check parent/vllm/model_executor/models/
|
||||||
|
candidates.append(os.path.join(p, "vllm", "model_executor", "models"))
|
||||||
|
|
||||||
|
# Repo-relative prebuilt bundle
|
||||||
|
here = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
candidates.append(os.path.join(here, "..", "..", "qwen3_6_scripts",
|
||||||
|
"prebuilt", "corex-3.2.3-ivcore10"))
|
||||||
|
candidates.append(os.path.join(here, "..", "..", "qwen3_6_scripts"))
|
||||||
|
|
||||||
|
# System CoreX
|
||||||
|
candidates.append("/usr/local/corex/lib64/")
|
||||||
|
|
||||||
|
for d in candidates:
|
||||||
|
d = os.path.normpath(d)
|
||||||
|
if os.path.isdir(d):
|
||||||
|
test_so = os.path.join(d, "corex_gdn_causal_conv.so")
|
||||||
|
if os.path.isfile(test_so):
|
||||||
|
return d
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _load_so(name: str, so_dir: str):
|
||||||
|
"""Load a single .so by name from so_dir via importlib."""
|
||||||
|
so_path = os.path.join(so_dir, f"{name}.so")
|
||||||
|
if not os.path.isfile(so_path):
|
||||||
|
return None
|
||||||
|
try:
|
||||||
|
spec = importlib.util.spec_from_file_location(name, so_path)
|
||||||
|
mod = importlib.util.module_from_spec(spec)
|
||||||
|
spec.loader.exec_module(mod)
|
||||||
|
return mod
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Failed to load %s: %s", so_path, e)
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
class CoreXModules:
|
||||||
|
"""Container for all loaded CoreX .so modules.
|
||||||
|
|
||||||
|
Each attribute is either the loaded module or None.
|
||||||
|
Attribute names drop the 'corex_' prefix for brevity.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._loaded = {}
|
||||||
|
self._so_dir = None
|
||||||
|
|
||||||
|
so_dir = _find_so_dir()
|
||||||
|
if so_dir is None:
|
||||||
|
logger.info("CoreX prebuilt .so directory not found — all modules disabled")
|
||||||
|
for name in _SO_MANIFEST:
|
||||||
|
short = name.replace("corex_", "", 1)
|
||||||
|
setattr(self, short, None)
|
||||||
|
self._loaded[name] = False
|
||||||
|
return
|
||||||
|
|
||||||
|
self._so_dir = so_dir
|
||||||
|
logger.info("CoreX .so directory: %s", so_dir)
|
||||||
|
|
||||||
|
loaded_count = 0
|
||||||
|
for name in _SO_MANIFEST:
|
||||||
|
mod = _load_so(name, so_dir)
|
||||||
|
short = name.replace("corex_", "", 1)
|
||||||
|
setattr(self, short, mod)
|
||||||
|
self._loaded[name] = mod is not None
|
||||||
|
if mod is not None:
|
||||||
|
loaded_count += 1
|
||||||
|
|
||||||
|
logger.info("CoreX: %d/%d .so loaded from %s",
|
||||||
|
loaded_count, len(_SO_MANIFEST), so_dir)
|
||||||
|
|
||||||
|
def summary(self) -> str:
|
||||||
|
"""Return a human-readable summary of loaded modules."""
|
||||||
|
lines = [f"CoreX .so loader ({self._so_dir or 'NOT FOUND'})"]
|
||||||
|
for name in _SO_MANIFEST:
|
||||||
|
status = "✓" if self._loaded.get(name) else "✗"
|
||||||
|
short = name.replace("corex_", "", 1)
|
||||||
|
mod = getattr(self, short, None)
|
||||||
|
if mod is not None:
|
||||||
|
funcs = [f for f in dir(mod) if not f.startswith("_")]
|
||||||
|
lines.append(f" {status} {name} → .{', .'.join(funcs)}")
|
||||||
|
else:
|
||||||
|
lines.append(f" {status} {name}")
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def all_loaded(self) -> bool:
|
||||||
|
return all(self._loaded.values())
|
||||||
|
|
||||||
|
@property
|
||||||
|
def loaded_count(self) -> int:
|
||||||
|
return sum(1 for v in self._loaded.values() if v)
|
||||||
|
|
||||||
|
|
||||||
|
# Singleton — initialized on first import
|
||||||
|
corex = CoreXModules()
|
||||||
343
vllm/ix_unified.py
Normal file
343
vllm/ix_unified.py
Normal file
@@ -0,0 +1,343 @@
|
|||||||
|
"""ix_unified.py — Unified Python interface to all ixformer::infer APIs.
|
||||||
|
|
||||||
|
Dispatch hierarchy (CCCL policy_selector pattern):
|
||||||
|
Tier 0: ix_unified_bridge.so (C++ direct call to ixformer::infer)
|
||||||
|
Tier 1: ixformer.functions.* (base image Python bindings, partial)
|
||||||
|
Tier 2: PyTorch fallback (always works, slowest)
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
from ex_engine.python.ix_unified import ix
|
||||||
|
out = ix.silu_and_mul(input)
|
||||||
|
ix.rms_norm(output, input, weight, eps)
|
||||||
|
weights, indices = ix.moe_topk_softmax(gating, topk, renorm)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import os
|
||||||
|
import sys
|
||||||
|
import importlib
|
||||||
|
import importlib.util
|
||||||
|
import torch
|
||||||
|
import logging
|
||||||
|
|
||||||
|
logger = logging.getLogger("ix_unified")
|
||||||
|
|
||||||
|
_bridge = None
|
||||||
|
|
||||||
|
|
||||||
|
def _load_bridge():
|
||||||
|
"""Load ix_unified_bridge.so from known locations."""
|
||||||
|
global _bridge
|
||||||
|
if _bridge is not None:
|
||||||
|
return _bridge
|
||||||
|
|
||||||
|
# Pre-load ixformer .so symbols into GLOBAL symbol table.
|
||||||
|
# ix_unified_bridge.so has undefined ixformer::infer::* symbols that get
|
||||||
|
# resolved at runtime. Python default import uses RTLD_LOCAL, so we must
|
||||||
|
# force RTLD_GLOBAL on the ixformer .so files BEFORE loading our bridge.
|
||||||
|
try:
|
||||||
|
import ctypes
|
||||||
|
|
||||||
|
# Phase 0: Load torch core libs first — ixformer depends on libc10.so etc.
|
||||||
|
try:
|
||||||
|
import torch as _torch
|
||||||
|
_torch_lib = os.path.join(os.path.dirname(_torch.__file__), "lib")
|
||||||
|
for _name in ["libc10.so", "libtorch_cpu.so", "libtorch.so",
|
||||||
|
"libc10_cuda.so", "libtorch_cuda.so", "libtorch_python.so"]:
|
||||||
|
_p = os.path.join(_torch_lib, _name)
|
||||||
|
if os.path.isfile(_p):
|
||||||
|
try:
|
||||||
|
ctypes.CDLL(_p, mode=ctypes.RTLD_GLOBAL)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Phase 1: libixformer.so (CUDA kernels)
|
||||||
|
# Phase 2: _ixformer_torch.so (torch extension with ixformer_torch_ext::*)
|
||||||
|
# ONLY these two — do NOT recursively load unknown .so (causes segfault)
|
||||||
|
_ixf_base = "/usr/local/corex/lib64/python3/dist-packages/ixformer"
|
||||||
|
if os.path.isdir(_ixf_base):
|
||||||
|
for _name in ["libixformer.so",
|
||||||
|
"_ixformer_torch.cpython-310-x86_64-linux-gnu.so"]:
|
||||||
|
_p = os.path.join(_ixf_base, _name)
|
||||||
|
if os.path.isfile(_p):
|
||||||
|
try:
|
||||||
|
ctypes.CDLL(_p, mode=ctypes.RTLD_GLOBAL)
|
||||||
|
logger.info("Preloaded: %s", _name)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
search_paths = []
|
||||||
|
|
||||||
|
# 1. Same directory as this file
|
||||||
|
here = os.path.dirname(os.path.abspath(__file__))
|
||||||
|
search_paths.append(os.path.join(here, "..", "build"))
|
||||||
|
search_paths.append(here)
|
||||||
|
|
||||||
|
# 2. Workspace build dirs (Docker / real machine)
|
||||||
|
search_paths.append("/workspace/ex_engine/build")
|
||||||
|
search_paths.append("/home/dylan/project_6/ex_engine/build")
|
||||||
|
|
||||||
|
# 2. vllm install root (where prebuilt .so are deployed)
|
||||||
|
for p in sys.path:
|
||||||
|
if "vllm" in p or "dist-packages" in p:
|
||||||
|
search_paths.append(p)
|
||||||
|
|
||||||
|
# 3. Explicit env var
|
||||||
|
env_path = os.getenv("IX_BRIDGE_PATH")
|
||||||
|
if env_path:
|
||||||
|
search_paths.insert(0, env_path)
|
||||||
|
|
||||||
|
for search_dir in search_paths:
|
||||||
|
for name in ["ix_unified_bridge.so",
|
||||||
|
"ix_unified_bridge.cpython-310-x86_64-linux-gnu.so"]:
|
||||||
|
so_path = os.path.join(search_dir, name)
|
||||||
|
if os.path.isfile(so_path):
|
||||||
|
try:
|
||||||
|
spec = importlib.util.spec_from_file_location(
|
||||||
|
"ix_unified_bridge", so_path)
|
||||||
|
mod = importlib.util.module_from_spec(spec)
|
||||||
|
spec.loader.exec_module(mod)
|
||||||
|
_bridge = mod
|
||||||
|
logger.info("ix_unified_bridge loaded from %s", so_path)
|
||||||
|
return _bridge
|
||||||
|
except (ImportError, OSError, SystemError) as e:
|
||||||
|
logger.warning("Bridge load failed (expected if ixformer "
|
||||||
|
"namespace mismatch): %s: %s",
|
||||||
|
os.path.basename(so_path), e)
|
||||||
|
continue
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Bridge load unexpected error: %s", e)
|
||||||
|
continue
|
||||||
|
|
||||||
|
logger.info("ix_unified_bridge.so not found, using fallback dispatch")
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
def _try_ixformer_functions():
|
||||||
|
"""Try importing ixformer.functions from base image."""
|
||||||
|
try:
|
||||||
|
import ixformer.functions as ixf
|
||||||
|
return ixf
|
||||||
|
except (ImportError, AttributeError):
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
# ============================================================================
|
||||||
|
# Dispatch class
|
||||||
|
# ============================================================================
|
||||||
|
|
||||||
|
class IXDispatch:
|
||||||
|
"""Three-tier dispatch for all ixformer ops."""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._bridge = _load_bridge()
|
||||||
|
self._ixf = _try_ixformer_functions()
|
||||||
|
tier = ("Tier0:bridge" if self._bridge else
|
||||||
|
"Tier1:ixformer" if self._ixf else "Tier2:pytorch")
|
||||||
|
logger.info("IXDispatch initialized: %s", tier)
|
||||||
|
|
||||||
|
# --- Activation -----------------------------------------------------------
|
||||||
|
def silu_and_mul(self, input: torch.Tensor) -> torch.Tensor:
|
||||||
|
if self._bridge:
|
||||||
|
return self._bridge.silu_and_mul(input)
|
||||||
|
if self._ixf and hasattr(self._ixf, 'silu_and_mul'):
|
||||||
|
d = input.size(-1) // 2
|
||||||
|
out = input.new_empty([input.size(0), d])
|
||||||
|
self._ixf.silu_and_mul(input, out)
|
||||||
|
return out
|
||||||
|
# PyTorch fallback
|
||||||
|
d = input.size(-1) // 2
|
||||||
|
x, gate = input[..., :d], input[..., d:]
|
||||||
|
return x * torch.sigmoid(gate)
|
||||||
|
|
||||||
|
# --- Norm -----------------------------------------------------------------
|
||||||
|
def rms_norm(self, output: torch.Tensor, input: torch.Tensor,
|
||||||
|
weight: torch.Tensor, eps: float):
|
||||||
|
if self._bridge:
|
||||||
|
self._bridge.rms_norm(output, input, weight, eps)
|
||||||
|
return
|
||||||
|
if self._ixf and hasattr(self._ixf, 'rms_norm'):
|
||||||
|
self._ixf.rms_norm(input, weight, output, eps)
|
||||||
|
return
|
||||||
|
# PyTorch fallback
|
||||||
|
variance = input.float().pow(2).mean(-1, keepdim=True)
|
||||||
|
normed = input * torch.rsqrt(variance + eps)
|
||||||
|
output.copy_(normed * weight)
|
||||||
|
|
||||||
|
def fused_add_rms_norm(self, input: torch.Tensor,
|
||||||
|
residual: torch.Tensor,
|
||||||
|
weight: torch.Tensor, eps: float):
|
||||||
|
if self._bridge:
|
||||||
|
self._bridge.fused_add_rms_norm(input, residual, weight, eps)
|
||||||
|
return
|
||||||
|
if self._ixf and hasattr(self._ixf, 'fused_add_rms_norm'):
|
||||||
|
self._ixf.fused_add_rms_norm(input, residual, weight, eps, 1.0)
|
||||||
|
return
|
||||||
|
# PyTorch fallback
|
||||||
|
hidden = input + residual
|
||||||
|
residual.copy_(hidden)
|
||||||
|
variance = hidden.float().pow(2).mean(-1, keepdim=True)
|
||||||
|
normed = hidden * torch.rsqrt(variance + eps)
|
||||||
|
input.copy_(normed * weight)
|
||||||
|
|
||||||
|
# --- Linear ---------------------------------------------------------------
|
||||||
|
def linear(self, input: torch.Tensor, weight: torch.Tensor,
|
||||||
|
bias=None) -> torch.Tensor:
|
||||||
|
if self._bridge:
|
||||||
|
return self._bridge.linear(input, weight, bias)
|
||||||
|
# PyTorch fallback
|
||||||
|
out = torch.nn.functional.linear(input, weight, bias)
|
||||||
|
return out
|
||||||
|
|
||||||
|
# --- RoPE -----------------------------------------------------------------
|
||||||
|
def rotary_embedding(self, positions, query, key, head_size,
|
||||||
|
cos_sin_cache, is_neox=True):
|
||||||
|
if self._bridge:
|
||||||
|
self._bridge.rotary_embedding(positions, query, key, head_size,
|
||||||
|
cos_sin_cache, is_neox)
|
||||||
|
return
|
||||||
|
if self._ixf and hasattr(self._ixf, 'vllm_rotary_embedding_neox'):
|
||||||
|
self._ixf.vllm_rotary_embedding_neox(
|
||||||
|
positions, query, key, head_size, cos_sin_cache, is_neox)
|
||||||
|
return
|
||||||
|
# No PyTorch fallback — this is handled by vllm's own rope
|
||||||
|
|
||||||
|
# --- KV Cache -------------------------------------------------------------
|
||||||
|
def reshape_and_cache(self, key, value, key_cache, value_cache,
|
||||||
|
slot_mapping):
|
||||||
|
if self._bridge:
|
||||||
|
self._bridge.reshape_and_cache(key, value, key_cache, value_cache,
|
||||||
|
slot_mapping)
|
||||||
|
return
|
||||||
|
if self._ixf and hasattr(self._ixf, 'vllm_cache_ops_reshape_and_cache'):
|
||||||
|
self._ixf.vllm_cache_ops_reshape_and_cache(
|
||||||
|
key, value, key_cache, value_cache, slot_mapping)
|
||||||
|
return
|
||||||
|
# PyTorch fallback — slot-by-slot copy
|
||||||
|
for i, slot in enumerate(slot_mapping):
|
||||||
|
if slot < 0:
|
||||||
|
continue
|
||||||
|
block_idx = slot // key_cache.size(2)
|
||||||
|
block_off = slot % key_cache.size(2)
|
||||||
|
key_cache[block_idx, :, block_off, :] = key[i]
|
||||||
|
value_cache[block_idx, :, block_off, :] = value[i]
|
||||||
|
|
||||||
|
# --- Attention: prefill ---------------------------------------------------
|
||||||
|
def flash_attn_prefill(self, query, key_cache, value_cache, output,
|
||||||
|
block_tables, cu_seq_q, cu_seq_k,
|
||||||
|
max_seq_q, max_seq_k, is_causal, scale):
|
||||||
|
if self._bridge:
|
||||||
|
return self._bridge.flash_attn_prefill(
|
||||||
|
query, key_cache, value_cache, output, block_tables,
|
||||||
|
cu_seq_q, cu_seq_k, max_seq_q, max_seq_k, is_causal, scale)
|
||||||
|
if self._ixf and hasattr(self._ixf, 'ixinfer_flash_attn_unpad'):
|
||||||
|
return self._ixf.ixinfer_flash_attn_unpad(
|
||||||
|
query, key_cache, value_cache, output, block_tables,
|
||||||
|
cu_seq_q, cu_seq_k, max_seq_q, max_seq_k,
|
||||||
|
is_causal, -1, -1, scale, 0.0, False, None, None, None)
|
||||||
|
raise RuntimeError("flash_attn_prefill: no backend available")
|
||||||
|
|
||||||
|
# --- Attention: decode (paged) -------------------------------------------
|
||||||
|
def paged_attention(self, output, query, key_cache, value_cache,
|
||||||
|
num_kv_heads, scale, block_tables, context_lens,
|
||||||
|
block_size, max_context_len):
|
||||||
|
if self._bridge:
|
||||||
|
return self._bridge.paged_attention(
|
||||||
|
output, query, key_cache, value_cache,
|
||||||
|
num_kv_heads, scale, block_tables, context_lens,
|
||||||
|
block_size, max_context_len)
|
||||||
|
if self._ixf and hasattr(self._ixf,
|
||||||
|
'vllm_single_query_cached_kv_attention_v2'):
|
||||||
|
return self._ixf.vllm_single_query_cached_kv_attention_v2(
|
||||||
|
output, query, key_cache, value_cache,
|
||||||
|
num_kv_heads, scale, block_tables, context_lens,
|
||||||
|
block_size, max_context_len, None)
|
||||||
|
raise RuntimeError("paged_attention: no backend available")
|
||||||
|
|
||||||
|
# --- MoE: topk_softmax ---------------------------------------------------
|
||||||
|
def moe_topk_softmax(self, gating_output: torch.Tensor,
|
||||||
|
topk: int, renormalize: bool = True):
|
||||||
|
if self._bridge:
|
||||||
|
return self._bridge.moe_topk_softmax(
|
||||||
|
gating_output, topk, renormalize)
|
||||||
|
# PyTorch fallback
|
||||||
|
scores = torch.softmax(gating_output.float(), dim=-1)
|
||||||
|
topk_weights, topk_indices = torch.topk(scores, k=topk, dim=-1)
|
||||||
|
if renormalize:
|
||||||
|
topk_weights = topk_weights / topk_weights.sum(dim=-1,
|
||||||
|
keepdim=True)
|
||||||
|
return topk_weights, topk_indices.to(torch.int32)
|
||||||
|
|
||||||
|
# --- MoE: gen_idx ---------------------------------------------------------
|
||||||
|
def moe_gen_idx(self, expert_ids: torch.Tensor, num_experts: int):
|
||||||
|
if self._bridge:
|
||||||
|
return self._bridge.moe_gen_idx(expert_ids, num_experts)
|
||||||
|
# PyTorch fallback: compute scatter/gather indices
|
||||||
|
flat = expert_ids.view(-1)
|
||||||
|
n = flat.numel()
|
||||||
|
src_dst = torch.empty(n, dtype=flat.dtype, device=flat.device)
|
||||||
|
dst_src = torch.empty(n, dtype=flat.dtype, device=flat.device)
|
||||||
|
expert_sizes = torch.zeros(num_experts, dtype=flat.dtype,
|
||||||
|
device=flat.device)
|
||||||
|
# Simple counting sort
|
||||||
|
for i in range(n):
|
||||||
|
expert_sizes[flat[i].item()] += 1
|
||||||
|
cumsum = expert_sizes.cumsum(-1)
|
||||||
|
offsets = torch.zeros_like(expert_sizes)
|
||||||
|
offsets[1:] = cumsum[:-1]
|
||||||
|
counts = torch.zeros_like(expert_sizes)
|
||||||
|
for i in range(n):
|
||||||
|
e = flat[i].item()
|
||||||
|
pos = (offsets[e] + counts[e]).item()
|
||||||
|
src_dst[i] = pos
|
||||||
|
dst_src[pos] = i
|
||||||
|
counts[e] += 1
|
||||||
|
return [src_dst, dst_src, expert_sizes, cumsum]
|
||||||
|
|
||||||
|
# --- MoE: expand_input ----------------------------------------------------
|
||||||
|
def moe_expand_input(self, input: torch.Tensor,
|
||||||
|
gather_index: torch.Tensor,
|
||||||
|
combine_idx: torch.Tensor, topk: int):
|
||||||
|
if self._bridge:
|
||||||
|
return self._bridge.moe_expand_input(
|
||||||
|
input, gather_index, combine_idx, topk)
|
||||||
|
# PyTorch fallback
|
||||||
|
return input.index_select(0, combine_idx.view(-1).long())
|
||||||
|
|
||||||
|
# --- MoE: group_gemm -----------------------------------------------------
|
||||||
|
def moe_group_gemm(self, input: torch.Tensor, weight: torch.Tensor,
|
||||||
|
tokens_per_experts: torch.Tensor):
|
||||||
|
if self._bridge:
|
||||||
|
return self._bridge.moe_group_gemm(
|
||||||
|
input, weight, tokens_per_experts)
|
||||||
|
# PyTorch fallback: sequential per-expert GEMM
|
||||||
|
outputs = []
|
||||||
|
offset = 0
|
||||||
|
for e in range(tokens_per_experts.size(0)):
|
||||||
|
count = tokens_per_experts[e].item()
|
||||||
|
if count == 0:
|
||||||
|
continue
|
||||||
|
inp_e = input[offset:offset + count]
|
||||||
|
w_e = weight[e] # [out_features, in_features]
|
||||||
|
outputs.append(inp_e @ w_e.t())
|
||||||
|
offset += count
|
||||||
|
if outputs:
|
||||||
|
return torch.cat(outputs, dim=0)
|
||||||
|
return input.new_empty(0, weight.size(-2))
|
||||||
|
|
||||||
|
# --- MoE: combine_result -------------------------------------------------
|
||||||
|
def moe_combine_result(self, expert_output: torch.Tensor,
|
||||||
|
weights: torch.Tensor):
|
||||||
|
if self._bridge:
|
||||||
|
return self._bridge.moe_combine_result(expert_output, weights)
|
||||||
|
# PyTorch fallback: weighted sum
|
||||||
|
# expert_output: [n_tokens, topk, hidden]
|
||||||
|
# weights: [n_tokens, topk]
|
||||||
|
return (expert_output * weights.unsqueeze(-1)).sum(dim=1)
|
||||||
|
|
||||||
|
|
||||||
|
# Singleton
|
||||||
|
ix = IXDispatch()
|
||||||
BIN
vllm/ix_unified_bridge.so
Executable file
BIN
vllm/ix_unified_bridge.so
Executable file
Binary file not shown.
236
vllm/moe_fused_dispatch.py
Normal file
236
vllm/moe_fused_dispatch.py
Normal file
@@ -0,0 +1,236 @@
|
|||||||
|
"""moe_fused_dispatch.py — Three-tier MoE dispatch (CCCL policy_selector pattern).
|
||||||
|
|
||||||
|
Port of upstream_ref/xllm/core/layers/ilu/fused_moe.cpp 7-step pipeline.
|
||||||
|
|
||||||
|
Dispatch hierarchy:
|
||||||
|
Tier 0: ix_unified_bridge.so → ixformer::infer 7-step C++ pipeline
|
||||||
|
topk_softmax → gen_idx → expand_input → group_gemm(w13) →
|
||||||
|
silu_and_mul → group_gemm(w2) → combine_result
|
||||||
|
Tier 1: corex prebuilt .so → direct_routed.w13/.w2_reduce (decode T=1 only)
|
||||||
|
Tier 2: PyTorch fallback → per-expert F.linear loop
|
||||||
|
|
||||||
|
Usage in qwen3_5.py:
|
||||||
|
from ex_engine.python.moe_fused_dispatch import fused_moe_forward
|
||||||
|
out = fused_moe_forward(hidden_states, router_logits, w13, w2,
|
||||||
|
top_k=8, num_experts=256, act_fn=silu_and_mul)
|
||||||
|
"""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
from typing import Callable, Optional
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
logger = logging.getLogger("moe_fused_dispatch")
|
||||||
|
|
||||||
|
# Lazy imports — set at first call
|
||||||
|
_ix = None
|
||||||
|
_corex = None
|
||||||
|
_init_done = False
|
||||||
|
|
||||||
|
|
||||||
|
def _lazy_init():
|
||||||
|
global _ix, _corex, _init_done
|
||||||
|
if _init_done:
|
||||||
|
return
|
||||||
|
_init_done = True
|
||||||
|
|
||||||
|
# Tier 0: ix_unified
|
||||||
|
try:
|
||||||
|
from ex_engine.python.ix_unified import ix
|
||||||
|
if ix._bridge is not None:
|
||||||
|
_ix = ix
|
||||||
|
logger.info("moe_fused_dispatch: Tier0 ix_unified_bridge.so available")
|
||||||
|
else:
|
||||||
|
logger.info("moe_fused_dispatch: Tier0 unavailable (bridge=None)")
|
||||||
|
except Exception as e:
|
||||||
|
logger.info("moe_fused_dispatch: Tier0 unavailable (%s)", e)
|
||||||
|
|
||||||
|
# Try import path used on real hardware
|
||||||
|
if _ix is None:
|
||||||
|
try:
|
||||||
|
from ix_unified import ix
|
||||||
|
if ix._bridge is not None:
|
||||||
|
_ix = ix
|
||||||
|
logger.info("moe_fused_dispatch: Tier0 ix_unified (direct) available")
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Tier 1: corex prebuilt .so
|
||||||
|
try:
|
||||||
|
from ex_engine.python.corex_so_loader import corex
|
||||||
|
if corex.moe_direct_routed is not None:
|
||||||
|
_corex = corex
|
||||||
|
logger.info("moe_fused_dispatch: Tier1 corex prebuilt .so available")
|
||||||
|
except Exception as e:
|
||||||
|
logger.info("moe_fused_dispatch: Tier1 unavailable (%s)", e)
|
||||||
|
|
||||||
|
|
||||||
|
def _tier0_fused_moe(
|
||||||
|
hidden_states: torch.Tensor, # [T, H]
|
||||||
|
router_logits: torch.Tensor, # [T, E]
|
||||||
|
w13: torch.Tensor, # [E, 2*I, H]
|
||||||
|
w2: torch.Tensor, # [E, H, I]
|
||||||
|
top_k: int,
|
||||||
|
num_experts: int,
|
||||||
|
act_fn: Callable,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Tier 0: Full 7-step ixformer::infer pipeline via ix_unified_bridge.so.
|
||||||
|
|
||||||
|
Maps 1:1 to xllm/core/layers/ilu/fused_moe.cpp::forward().
|
||||||
|
"""
|
||||||
|
T, H = hidden_states.shape
|
||||||
|
|
||||||
|
# Step 1: topk_softmax — fused softmax + topk selection
|
||||||
|
topk_weights, topk_ids = _ix.moe_topk_softmax(router_logits, top_k,
|
||||||
|
renormalize=True)
|
||||||
|
|
||||||
|
# Step 2: gen_idx — compute scatter/gather indices for expert routing
|
||||||
|
idx_result = _ix.moe_gen_idx(topk_ids, num_experts)
|
||||||
|
src_dst, dst_src, expert_sizes, cumsum = idx_result
|
||||||
|
|
||||||
|
# Step 3: expand_input — scatter tokens to expert order
|
||||||
|
expanded = _ix.moe_expand_input(hidden_states, dst_src, src_dst, top_k)
|
||||||
|
|
||||||
|
# Step 4: group_gemm(w13) — batched GEMM across all experts
|
||||||
|
gate_up = _ix.moe_group_gemm(expanded, w13, expert_sizes)
|
||||||
|
|
||||||
|
# Step 5: activation — SiLU(gate) * up
|
||||||
|
act = act_fn(gate_up)
|
||||||
|
|
||||||
|
# Step 6: group_gemm(w2) — down projection
|
||||||
|
down = _ix.moe_group_gemm(act, w2, expert_sizes)
|
||||||
|
|
||||||
|
# Step 7: combine_result — gather back and weighted sum
|
||||||
|
output = _ix.moe_combine_result(
|
||||||
|
down.view(T, top_k, H), topk_weights)
|
||||||
|
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
|
def _tier1_decode_single_token(
|
||||||
|
hidden_states: torch.Tensor, # [1, H]
|
||||||
|
expert_ids: torch.Tensor, # [K]
|
||||||
|
weights: torch.Tensor, # [K]
|
||||||
|
w13: torch.Tensor, # [E, 2*I, H]
|
||||||
|
w2: torch.Tensor, # [E, H, I]
|
||||||
|
act_fn: Callable,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Tier 1: Single-token decode via prebuilt corex_moe_direct_routed.so.
|
||||||
|
|
||||||
|
Only works for T=1 decode. The .so implements fused expert indexing +
|
||||||
|
GEMM + reduction in a single kernel launch.
|
||||||
|
"""
|
||||||
|
gate_up = _corex.moe_direct_routed.w13(hidden_states, w13, expert_ids)
|
||||||
|
act = act_fn(gate_up)
|
||||||
|
return _corex.moe_direct_routed.w2_reduce(act, w2, expert_ids, weights)
|
||||||
|
|
||||||
|
|
||||||
|
def _tier2_pytorch_loop(
|
||||||
|
hidden_states: torch.Tensor, # [T, H]
|
||||||
|
router_logits: torch.Tensor, # [T, E]
|
||||||
|
w13: torch.Tensor, # [E, 2*I, H]
|
||||||
|
w2: torch.Tensor, # [E, H, I]
|
||||||
|
top_k: int,
|
||||||
|
act_fn: Callable,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Tier 2: Pure PyTorch per-expert loop (always works, slowest)."""
|
||||||
|
T, H = hidden_states.shape
|
||||||
|
|
||||||
|
# Softmax → topk
|
||||||
|
topk_logits, topk_ids = torch.topk(router_logits.float(), top_k, dim=-1)
|
||||||
|
topk_weights = torch.softmax(topk_logits, dim=-1).to(hidden_states.dtype)
|
||||||
|
|
||||||
|
if T == 1:
|
||||||
|
# Fast single-token path: batched GEMM
|
||||||
|
eids = topk_ids[0]
|
||||||
|
ws = topk_weights[0]
|
||||||
|
w13_sel = w13[eids]
|
||||||
|
w2_sel = w2[eids]
|
||||||
|
gate_up = F.linear(hidden_states, w13_sel.reshape(-1, H))
|
||||||
|
gate_up = gate_up.view(top_k, -1)
|
||||||
|
act = act_fn(gate_up)
|
||||||
|
expert_out = torch.bmm(w2_sel, act.unsqueeze(-1)).squeeze(-1)
|
||||||
|
return (expert_out * ws.unsqueeze(-1)).sum(0, keepdim=True).to(
|
||||||
|
hidden_states.dtype)
|
||||||
|
else:
|
||||||
|
# General prefill path: sorted per-expert loop
|
||||||
|
out = torch.zeros_like(hidden_states)
|
||||||
|
flat_eids = topk_ids.reshape(-1)
|
||||||
|
order = torch.argsort(flat_eids, stable=True)
|
||||||
|
sorted_tok_ids = torch.arange(
|
||||||
|
T, device=topk_ids.device).repeat_interleave(top_k)[order]
|
||||||
|
sorted_weights = topk_weights.reshape(-1)[order]
|
||||||
|
expert_counts = torch.bincount(
|
||||||
|
flat_eids, minlength=w13.shape[0]).tolist()
|
||||||
|
|
||||||
|
start = 0
|
||||||
|
for eid, count in enumerate(expert_counts):
|
||||||
|
if count == 0:
|
||||||
|
continue
|
||||||
|
end = start + count
|
||||||
|
tok_ids = sorted_tok_ids[start:end]
|
||||||
|
tokens = hidden_states[tok_ids]
|
||||||
|
gate_up = F.linear(tokens, w13[eid])
|
||||||
|
act = act_fn(gate_up)
|
||||||
|
expert_out = F.linear(act, w2[eid])
|
||||||
|
weights_e = sorted_weights[start:end].unsqueeze(-1)
|
||||||
|
out.index_add_(0, tok_ids, (expert_out * weights_e).to(out.dtype))
|
||||||
|
start = end
|
||||||
|
return out
|
||||||
|
|
||||||
|
|
||||||
|
def fused_moe_forward(
|
||||||
|
hidden_states: torch.Tensor, # [T, H]
|
||||||
|
router_logits: torch.Tensor, # [T, E]
|
||||||
|
w13: torch.Tensor, # [E, 2*I, H]
|
||||||
|
w2: torch.Tensor, # [E, H, I]
|
||||||
|
top_k: int = 8,
|
||||||
|
num_experts: int = 256,
|
||||||
|
act_fn: Optional[Callable] = None,
|
||||||
|
) -> torch.Tensor:
|
||||||
|
"""Dispatch MoE through Tier 0 → 1 → 2.
|
||||||
|
|
||||||
|
Returns partial output (pre all-reduce), same contract as vllm FusedMoE.
|
||||||
|
"""
|
||||||
|
_lazy_init()
|
||||||
|
|
||||||
|
if act_fn is None:
|
||||||
|
def _default_act(x):
|
||||||
|
gate, up = x.chunk(2, dim=-1)
|
||||||
|
return F.silu(gate) * up
|
||||||
|
act_fn = _default_act
|
||||||
|
|
||||||
|
T = hidden_states.shape[0]
|
||||||
|
|
||||||
|
# Tier 0: full ixformer pipeline (all sizes)
|
||||||
|
if _ix is not None and _ix._bridge is not None:
|
||||||
|
try:
|
||||||
|
return _tier0_fused_moe(hidden_states, router_logits, w13, w2,
|
||||||
|
top_k, num_experts, act_fn)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Tier0 MoE failed (%s), falling to Tier1/2", e)
|
||||||
|
|
||||||
|
# Tier 1: corex direct routed (decode T=1 only)
|
||||||
|
if (T == 1 and _corex is not None
|
||||||
|
and _corex.moe_direct_routed is not None
|
||||||
|
and hidden_states.dtype == torch.float16
|
||||||
|
and w13.dtype == torch.float16
|
||||||
|
and w2.dtype == torch.float16
|
||||||
|
and hidden_states.is_contiguous()
|
||||||
|
and w13.is_contiguous()
|
||||||
|
and w2.is_contiguous()):
|
||||||
|
try:
|
||||||
|
topk_logits, topk_ids = torch.topk(
|
||||||
|
router_logits.float(), top_k, dim=-1)
|
||||||
|
topk_weights = torch.softmax(topk_logits, dim=-1).to(
|
||||||
|
hidden_states.dtype)
|
||||||
|
return _tier1_decode_single_token(
|
||||||
|
hidden_states, topk_ids[0], topk_weights[0],
|
||||||
|
w13, w2, act_fn)
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning("Tier1 MoE failed (%s), falling to Tier2", e)
|
||||||
|
|
||||||
|
# Tier 2: PyTorch fallback
|
||||||
|
return _tier2_pytorch_loop(hidden_states, router_logits, w13, w2,
|
||||||
|
top_k, act_fn)
|
||||||
Reference in New Issue
Block a user