feat: import CUDA kernels from xllm/CCCL/FLA upstream repos

Sources cloned and tree'd (no --depth):
  - jd-opensource/xllm: ILU kernels, CUDA kernels, MoE kernels
  - NVIDIA/cccl: CUB tuning/dispatch headers (block-level primitives)
  - fla-org/flash-linear-attention: Triton GDN kernels
  - NVIDIA/cutlass: grouped GEMM reference (read, not copied)
  - Dao-AILab/flash-attention: attention kernel reference (SM80+, read only)

New CUDA kernels (from xllm, SM-agnostic, portable to BI-V100):
  ex_engine/xllm_kernels/cuda/activation.cu    (188 lines) — silu_and_mul, gelu
  ex_engine/xllm_kernels/cuda/norm.cu          (600 lines) — rms_norm, fused_add_rms_norm
  ex_engine/xllm_kernels/cuda/rope.cu          (258 lines) — rotary_embedding
  ex_engine/xllm_kernels/cuda/block_copy.cu    (209 lines) — copy_blocks, swap_blocks
  ex_engine/xllm_kernels/cuda/reshape_paged_cache.cu (101 lines) — KV cache ops
  ex_engine/xllm_kernels/cuda/headers/         (5 headers for compilation)

ILU bridge kernel sources (from xllm, verified SAME as upstream):
  ex_engine/xllm_kernels/ilu/    (10 files, 925 lines total)
  — activation.cpp, attention.cpp, fused_moe.cpp, group_gemm.cpp,
    matmul.cpp, norm.cpp, rope.cpp, ilu_ops_api.h, ixformer.h, utils.h

FLA Triton GDN kernels (for GatedDeltaNet without SM90+ FlashQLA):
  ex_engine/fla_kernels/gated_delta_rule/  (7 files, 2370 lines)
  — chunk_fwd.py (428), chunk.py (487), wy_fast.py (409),
    fused_recurrent.py (392), naive.py (161), gate.py (380)

CCCL sync (12 tuning + 14 dispatch headers updated from NVIDIA/cccl):
  cccl_upstream/cub/cub/device/dispatch/tuning/ — 12 changed files synced
  cccl_upstream/cub/cub/device/dispatch/ — 14 changed dispatch files synced

Compilation targets for real machine (ivcore10):
  1. CUDA kernels: --cuda-gpu-arch=ivcore10 via corex clang/16
  2. ILU bridges: torch.utils.cpp_extension linking ixformer .so
  3. FLA kernels: Triton JIT (if Triton works on BI-V100)
This commit is contained in:
claude
2026-08-14 07:48:52 +00:00
parent 051b02d3cd
commit 8d75652949
66 changed files with 12834 additions and 5458 deletions

View File

@@ -0,0 +1,188 @@
/* Copyright 2025 The vLLM Authors and 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.
==============================================================================*/
#include <c10/cuda/CUDAGuard.h>
#include <torch/cuda.h>
#include <cstdint>
#include "cuda_ops_api.h"
#include "device_utils.cuh"
// ref to:
// https://github.com/vllm-project/vllm/blob/main/csrc/activation_kernels.cu
namespace {
using ::xllm::kernel::cuda::xllm_ldg;
template <typename scalar_t,
scalar_t (*ACT_FN)(const scalar_t&),
bool act_first>
__device__ __forceinline__ scalar_t compute(const scalar_t& x,
const scalar_t& y) {
return act_first ? ACT_FN(x) * y : x * ACT_FN(y);
}
// Check if pointer is 16-byte aligned for int4 vectorized access
__device__ __forceinline__ bool is_16byte_aligned(const void* ptr) {
return (reinterpret_cast<uintptr_t>(ptr) & 15) == 0;
}
// Activation and gating kernel template with 128-bit vectorized access
// optimization.
template <typename scalar_t,
scalar_t (*ACT_FN)(const scalar_t&),
bool act_first>
__global__ void XLLM_KERNEL_ATTR(1024)
act_and_mul_kernel(scalar_t* __restrict__ out, // [..., d]
const scalar_t* __restrict__ input, // [..., 2, d]
const int d) {
constexpr int kVecSize = 16 / sizeof(scalar_t);
const int64_t token_idx = blockIdx.x;
const scalar_t* x_ptr = input + token_idx * 2 * d;
const scalar_t* y_ptr = x_ptr + d;
scalar_t* out_ptr = out + token_idx * d;
// Check alignment for 128-bit vectorized access.
// All three pointers must be 16-byte aligned for safe int4 operations.
const bool aligned = is_16byte_aligned(x_ptr) && is_16byte_aligned(y_ptr) &&
is_16byte_aligned(out_ptr);
if (aligned && d >= kVecSize) {
// Fast path: 128-bit vectorized loop
const int4* x_vec = reinterpret_cast<const int4*>(x_ptr);
const int4* y_vec = reinterpret_cast<const int4*>(y_ptr);
int4* out_vec = reinterpret_cast<int4*>(out_ptr);
const int num_vecs = d / kVecSize;
const int vec_end = num_vecs * kVecSize;
for (int i = threadIdx.x; i < num_vecs; i += blockDim.x) {
int4 x = xllm_ldg(&x_vec[i]), y = xllm_ldg(&y_vec[i]), r;
auto* xp = reinterpret_cast<scalar_t*>(&x);
auto* yp = reinterpret_cast<scalar_t*>(&y);
auto* rp = reinterpret_cast<scalar_t*>(&r);
#pragma unroll
for (int j = 0; j < kVecSize; j++) {
rp[j] = compute<scalar_t, ACT_FN, act_first>(xp[j], yp[j]);
}
out_vec[i] = r;
}
// Scalar cleanup for remaining elements
for (int i = vec_end + threadIdx.x; i < d; i += blockDim.x) {
out_ptr[i] = compute<scalar_t, ACT_FN, act_first>(xllm_ldg(&x_ptr[i]),
xllm_ldg(&y_ptr[i]));
}
} else {
// Scalar fallback for unaligned data or small d
for (int64_t idx = threadIdx.x; idx < d; idx += blockDim.x) {
const scalar_t x = xllm_ldg(&x_ptr[idx]);
const scalar_t y = xllm_ldg(&y_ptr[idx]);
out_ptr[idx] = compute<scalar_t, ACT_FN, act_first>(x, y);
}
}
}
template <typename T>
__device__ __forceinline__ T silu_kernel(const T& x) {
// x * sigmoid(x)
const float f = static_cast<float>(x);
return static_cast<T>(f / (1.0f + expf(-f)));
}
template <typename T>
__device__ __forceinline__ T gelu_kernel(const T& x) {
// Equivalent to PyTorch GELU with 'none' approximation.
// Refer to:
// https://github.com/pytorch/pytorch/blob/8ac9b20d4b090c213799e81acf48a55ea8d437d6/aten/src/ATen/native/cuda/ActivationGeluKernel.cu#L36-L38
const float f = static_cast<float>(x);
constexpr float kAlpha = M_SQRT1_2;
return static_cast<T>(f * 0.5f * (1.0f + ::erf(f * kAlpha)));
}
template <typename T>
__device__ __forceinline__ T gelu_tanh_kernel(const T& x) {
// Equivalent to PyTorch GELU with 'tanh' approximation.
// Refer to:
// https://github.com/pytorch/pytorch/blob/8ac9b20d4b090c213799e81acf48a55ea8d437d6/aten/src/ATen/native/cuda/ActivationGeluKernel.cu#L25-L30
const float f = static_cast<float>(x);
constexpr float kBeta = M_SQRT2 * M_2_SQRTPI * 0.5f;
constexpr float kKappa = 0.044715;
float x_cube = f * f * f;
float inner = kBeta * (f + kKappa * x_cube);
return static_cast<T>(0.5f * f * (1.0f + ::tanhf(inner)));
}
#define LAUNCH_ACTIVATION_GATE_KERNEL(KERNEL, ACT_FIRST) \
int d = input.size(-1) / 2; \
int64_t num_tokens = input.numel() / input.size(-1); \
dim3 grid(num_tokens); \
dim3 block(std::min(d, 1024)); \
if (num_tokens == 0) { \
return; \
} \
const at::cuda::OptionalCUDAGuard device_guard(device_of(input)); \
const cudaStream_t stream = at::cuda::getCurrentCUDAStream(); \
DISPATCH_FLOATING_TYPES(input.scalar_type(), "act_and_mul_kernel", [&] { \
act_and_mul_kernel<scalar_t, KERNEL<scalar_t>, ACT_FIRST> \
<<<grid, block, 0, stream>>>( \
out.data_ptr<scalar_t>(), input.data_ptr<scalar_t>(), d); \
});
void silu_and_mul(torch::Tensor out, // [..., d]
torch::Tensor input) // [..., 2 * d]
{
LAUNCH_ACTIVATION_GATE_KERNEL(silu_kernel, true);
}
void gelu_and_mul(torch::Tensor& out, // [..., d]
torch::Tensor& input) // [..., 2 * d]
{
LAUNCH_ACTIVATION_GATE_KERNEL(gelu_kernel, true);
}
void gelu_tanh_and_mul(torch::Tensor& out, // [..., d]
torch::Tensor& input) // [..., 2 * d]
{
LAUNCH_ACTIVATION_GATE_KERNEL(gelu_tanh_kernel, true);
}
} // namespace
namespace xllm::kernel::cuda {
void act_and_mul(torch::Tensor out,
torch::Tensor input,
const std::string& act_mode) {
if (act_mode != "silu" && act_mode != "gelu" && act_mode != "gelu_tanh" &&
act_mode != "gelu_pytorch_tanh") {
LOG(FATAL) << "Unsupported act mode: " << act_mode
<< ", only support silu, gelu, gelu_tanh, gelu_pytorch_tanh";
}
// flashinfer act_and_mul ops
// std::string uri = act_mode + "_and_mul";
// FunctionFactory::get_instance().act_and_mul(uri).call(
// out, input, support_pdl());
if (act_mode == "silu") {
silu_and_mul(out, input);
} else if (act_mode == "gelu") {
gelu_and_mul(out, input);
} else if (act_mode == "gelu_tanh" || act_mode == "gelu_pytorch_tanh") {
// gelu_tanh or gelu_pytorch_tanh (mathematically equivalent)
gelu_tanh_and_mul(out, input);
}
}
} // namespace xllm::kernel::cuda

View File

@@ -0,0 +1,209 @@
/* Copyright 2025-2026 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.
==============================================================================*/
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAException.h>
#include <c10/cuda/CUDAGuard.h>
#include <c10/cuda/CUDAStream.h>
#include <cuda_runtime.h>
#include <torch/extension.h>
#include <cstdint>
#include <type_traits>
#include "cuda_ops_api.h"
#include "utils.h"
namespace xllm::kernel::cuda {
namespace {
template <typename scalar_t>
struct VecType;
template <>
struct VecType<c10::Half> {
using type = uint4;
static constexpr int32_t vec_width = 8;
};
template <>
struct VecType<c10::BFloat16> {
using type = uint4;
static constexpr int32_t vec_width = 8;
};
template <>
struct VecType<float> {
using type = float4;
static constexpr int32_t vec_width = 4;
};
DEVICE_INLINE int32_t find_group_idx(const int32_t* __restrict__ cum_sum,
const int32_t num_groups,
const int32_t dst_idx) {
int32_t left = 0;
int32_t right = num_groups - 1;
while (left < right) {
const int32_t mid = left + ((right - left) >> 1);
const bool move_left = dst_idx < cum_sum[mid];
right = move_left ? mid : right;
left = move_left ? left : mid + 1;
}
return left;
}
template <typename scalar_t, bool kVectorized>
__global__ void block_copy_kernel(const int64_t* __restrict__ key_cache_ptrs,
const int64_t* __restrict__ value_cache_ptrs,
const int32_t* __restrict__ src_block_indices,
const int32_t* __restrict__ dst_block_indices,
const int32_t* __restrict__ cum_sum,
const int32_t num_groups,
const int64_t numel_per_block) {
const int64_t layer_idx = static_cast<int64_t>(blockIdx.x);
const int32_t dst_linear_idx = static_cast<int32_t>(blockIdx.y);
const int64_t tile_idx = static_cast<int64_t>(blockIdx.z);
scalar_t* __restrict__ key_cache = reinterpret_cast<scalar_t*>(
static_cast<uintptr_t>(key_cache_ptrs[layer_idx]));
scalar_t* __restrict__ value_cache = reinterpret_cast<scalar_t*>(
static_cast<uintptr_t>(value_cache_ptrs[layer_idx]));
const int32_t group_idx = find_group_idx(cum_sum, num_groups, dst_linear_idx);
const int32_t src_block = src_block_indices[group_idx];
const int32_t dst_block = dst_block_indices[dst_linear_idx];
const int64_t src_offset = static_cast<int64_t>(src_block) * numel_per_block;
const int64_t dst_offset = static_cast<int64_t>(dst_block) * numel_per_block;
if constexpr (kVectorized) {
using VecTypeT = typename VecType<scalar_t>::type;
constexpr int32_t kVecWidth = VecType<scalar_t>::vec_width;
const int64_t num_vecs_per_block = numel_per_block / kVecWidth;
const int64_t vec_idx = tile_idx * static_cast<int64_t>(blockDim.x) +
static_cast<int64_t>(threadIdx.x);
if (vec_idx >= num_vecs_per_block) {
return;
}
const int64_t elem_offset = vec_idx * kVecWidth;
const auto* key_src_vec =
reinterpret_cast<const VecTypeT*>(key_cache + src_offset + elem_offset);
const auto* value_src_vec = reinterpret_cast<const VecTypeT*>(
value_cache + src_offset + elem_offset);
auto* key_dst_vec =
reinterpret_cast<VecTypeT*>(key_cache + dst_offset + elem_offset);
auto* value_dst_vec =
reinterpret_cast<VecTypeT*>(value_cache + dst_offset + elem_offset);
*key_dst_vec = *key_src_vec;
*value_dst_vec = *value_src_vec;
} else {
const int64_t elem_idx = tile_idx * static_cast<int64_t>(blockDim.x) +
static_cast<int64_t>(threadIdx.x);
if (elem_idx >= numel_per_block) {
return;
}
key_cache[dst_offset + elem_idx] = key_cache[src_offset + elem_idx];
value_cache[dst_offset + elem_idx] = value_cache[src_offset + elem_idx];
}
}
} // namespace
void block_copy(torch::Tensor key_cache_ptrs,
torch::Tensor value_cache_ptrs,
torch::Tensor src_block_indices,
torch::Tensor dst_block_indices,
torch::Tensor cum_sum,
int64_t numel_per_block,
torch::ScalarType cache_dtype) {
if (src_block_indices.numel() == 0) {
return;
}
CHECK(key_cache_ptrs.is_cuda());
CHECK(value_cache_ptrs.is_cuda());
CHECK(src_block_indices.is_cuda());
CHECK(dst_block_indices.is_cuda());
CHECK(cum_sum.is_cuda());
CHECK_EQ(key_cache_ptrs.scalar_type(), torch::kInt64);
CHECK_EQ(value_cache_ptrs.scalar_type(), torch::kInt64);
CHECK_EQ(src_block_indices.scalar_type(), torch::kInt32);
CHECK_EQ(dst_block_indices.scalar_type(), torch::kInt32);
CHECK_EQ(cum_sum.scalar_type(), torch::kInt32);
CHECK_EQ(key_cache_ptrs.dim(), 1);
CHECK_EQ(value_cache_ptrs.dim(), 1);
CHECK_EQ(src_block_indices.dim(), 1);
CHECK_EQ(dst_block_indices.dim(), 1);
CHECK_EQ(cum_sum.dim(), 1);
CHECK(key_cache_ptrs.is_contiguous());
CHECK(value_cache_ptrs.is_contiguous());
CHECK(src_block_indices.is_contiguous());
CHECK(dst_block_indices.is_contiguous());
CHECK(cum_sum.is_contiguous());
CHECK_EQ(key_cache_ptrs.size(0), value_cache_ptrs.size(0));
CHECK_EQ(src_block_indices.size(0), cum_sum.size(0));
CHECK_GT(numel_per_block, 0);
const at::cuda::OptionalCUDAGuard device_guard(key_cache_ptrs.device());
constexpr int32_t kThreadsPerBlock = 256;
const int32_t num_layers = static_cast<int32_t>(key_cache_ptrs.size(0));
const int32_t num_groups = static_cast<int32_t>(src_block_indices.size(0));
const int32_t num_dst_blocks =
static_cast<int32_t>(dst_block_indices.size(0));
const cudaStream_t stream =
c10::cuda::getCurrentCUDAStream(key_cache_ptrs.get_device());
DISPATCH_FLOATING_TYPES(cache_dtype, "block_copy_kernel", [&] {
constexpr bool kHasVecType = std::is_same_v<scalar_t, float> ||
std::is_same_v<scalar_t, c10::Half> ||
std::is_same_v<scalar_t, c10::BFloat16>;
if constexpr (kHasVecType) {
constexpr int32_t kVecWidth = VecType<scalar_t>::vec_width;
if (numel_per_block % kVecWidth == 0) {
const int64_t tiles_per_block =
ceil_div<int64_t>(numel_per_block / kVecWidth, kThreadsPerBlock);
const dim3 grid(num_layers, num_dst_blocks, tiles_per_block);
block_copy_kernel<scalar_t, true>
<<<grid, kThreadsPerBlock, 0, stream>>>(
key_cache_ptrs.data_ptr<int64_t>(),
value_cache_ptrs.data_ptr<int64_t>(),
src_block_indices.data_ptr<int32_t>(),
dst_block_indices.data_ptr<int32_t>(),
cum_sum.data_ptr<int32_t>(),
num_groups,
numel_per_block);
C10_CUDA_KERNEL_LAUNCH_CHECK();
return;
}
}
const int64_t tiles_per_block =
ceil_div<int64_t>(numel_per_block, kThreadsPerBlock);
const dim3 grid(num_layers, num_dst_blocks, tiles_per_block);
block_copy_kernel<scalar_t, false><<<grid, kThreadsPerBlock, 0, stream>>>(
key_cache_ptrs.data_ptr<int64_t>(),
value_cache_ptrs.data_ptr<int64_t>(),
src_block_indices.data_ptr<int32_t>(),
dst_block_indices.data_ptr<int32_t>(),
cum_sum.data_ptr<int32_t>(),
num_groups,
numel_per_block);
C10_CUDA_KERNEL_LAUNCH_CHECK();
});
}
} // namespace xllm::kernel::cuda

View File

@@ -0,0 +1,306 @@
/* Copyright 2025-2026 The xLLM Authors.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://github.com/jd-opensource/xllm/blob/main/LICENSE
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include <ATen/DynamicLibrary.h>
#include <ATen/core/dispatch/Dispatcher.h>
#include <glog/logging.h>
#include <optional>
#include <tuple>
#include <vector>
#include "utils.h"
namespace xllm::kernel::cuda {
// TODO: add head_size parameter
void rotary_embedding(torch::Tensor& positions,
torch::Tensor& query,
std::optional<torch::Tensor> key,
torch::Tensor& cos_sin_cache,
// int64_t head_size,
bool is_neox);
// act_mode only support silu, gelu, gelu_tanh
void act_and_mul(torch::Tensor out,
torch::Tensor input,
const std::string& act_mode);
void reshape_paged_cache(
torch::Tensor slot_ids, // [n_tokens]
torch::Tensor keys, // [n_tokens, n_kv_heads, head_dim]
torch::Tensor values, // [n_tokens, n_kv_heads, head_dim]
torch::Tensor key_cache, // [n_blocks, block_size, n_heads, head_dim]
torch::Tensor value_cache);
void block_copy(torch::Tensor key_cache_ptrs,
torch::Tensor value_cache_ptrs,
torch::Tensor src_block_indices,
torch::Tensor dst_block_indices,
torch::Tensor cum_sum,
int64_t numel_per_block,
torch::ScalarType cache_dtype);
#if !defined(USE_DCU)
void batch_prefill(const std::string& uri,
ffi::Array<int64_t> plan_info,
torch::Tensor float_workspace_buffer,
torch::Tensor int_workspace_buffer,
torch::Tensor page_locked_int_workspace_buffer,
torch::Tensor query,
torch::Tensor key,
torch::Tensor value,
torch::Tensor q_cu_seq_lens,
torch::Tensor kv_cu_seq_lens,
int64_t window_left,
double sm_scale,
torch::Tensor output,
std::optional<torch::Tensor>& output_lse,
const std::optional<torch::Tensor>& mask = std::nullopt);
// Wrapper function for batch_prefill that conditionally uses AttentionRunner
// for piecewise CUDA Graph capture
void batch_prefill_with_optional_piecewise_capture(
const std::string& uri,
ffi::Array<int64_t> plan_info,
torch::Tensor float_workspace_buffer,
torch::Tensor int_workspace_buffer,
torch::Tensor page_locked_int_workspace_buffer,
torch::Tensor query,
torch::Tensor key,
torch::Tensor value,
torch::Tensor q_cu_seq_lens,
torch::Tensor kv_cu_seq_lens,
int64_t window_left,
double sm_scale,
torch::Tensor output,
std::optional<torch::Tensor>& output_lse);
void batch_prefill_non_causal(
const std::string& uri,
ffi::Array<int64_t> plan_info,
torch::Tensor float_workspace_buffer,
torch::Tensor int_workspace_buffer,
torch::Tensor page_locked_int_workspace_buffer,
torch::Tensor query,
torch::Tensor key,
torch::Tensor value,
torch::Tensor q_cu_seq_lens,
torch::Tensor kv_cu_seq_lens,
int64_t window_left,
double sm_scale,
torch::Tensor output,
std::optional<torch::Tensor>& output_lse,
const std::optional<torch::Tensor>& mask = std::nullopt);
void batch_chunked_prefill(
const std::string& uri,
ffi::Array<int64_t> plan_info,
torch::Tensor float_workspace_buffer,
torch::Tensor int_workspace_buffer,
torch::Tensor page_locked_int_workspace_buffer,
torch::Tensor query,
torch::Tensor k_cache,
torch::Tensor v_cache,
torch::Tensor paged_kv_indptr,
torch::Tensor paged_kv_indices,
torch::Tensor paged_kv_last_page_len,
int64_t window_left,
double sm_scale,
torch::Tensor output,
std::optional<torch::Tensor>& output_lse,
std::optional<torch::Tensor> qo_indptr = std::nullopt,
bool causal = true);
void batch_decode(const std::string& uri,
ffi::Array<int64_t> plan_info,
torch::Tensor float_workspace_buffer,
torch::Tensor int_workspace_buffer,
torch::Tensor page_locked_int_workspace_buffer,
torch::Tensor query,
torch::Tensor k_cache,
torch::Tensor v_cache,
torch::Tensor paged_kv_indptr,
torch::Tensor paged_kv_indices,
torch::Tensor paged_kv_last_page_len,
int64_t window_left,
double sm_scale,
torch::Tensor output,
std::optional<torch::Tensor>& output_lse,
bool use_tensor_core,
std::optional<torch::Tensor> qo_indptr = std::nullopt);
#endif // !defined(USE_DCU)
void rms_norm(torch::Tensor output,
torch::Tensor input,
torch::Tensor weight,
double eps);
void fused_add_rms_norm(torch::Tensor& input, // [..., hidden_size]
torch::Tensor& residual, // [..., hidden_size]
torch::Tensor& weight, // [hidden_size]
double epsilon);
torch::Tensor matmul(torch::Tensor a,
torch::Tensor b,
std::optional<torch::Tensor> bias);
void cutlass_scaled_mm(torch::Tensor& c,
torch::Tensor const& a,
torch::Tensor const& b,
torch::Tensor const& a_scales,
torch::Tensor const& b_scales,
std::optional<torch::Tensor> const& bias);
// Static scaled FP8 quantization
// Quantizes input tensor to FP8 using a pre-computed scale factor
void static_scaled_fp8_quant(torch::Tensor& out, // [..., d]
torch::Tensor const& input, // [..., d]
torch::Tensor const& scale); // [1]
// FP8 scaled quantize: quantizes input tensor to FP8 e4m3 format
// Returns: (quantized_output, scale)
std::tuple<torch::Tensor, torch::Tensor> fp8_scaled_quantize(
const torch::Tensor& input,
const std::optional<torch::Tensor>& output = std::nullopt,
const std::optional<torch::Tensor>& scale = std::nullopt);
// ============================================================================
// Fused RMSNorm + Static FP8 Quantization
// ============================================================================
// These functions combine RMSNorm and FP8 quantization to reduce memory
// bandwidth by avoiding the intermediate write-back to global memory.
// Fused RMSNorm + Static FP8 Quantization (without residual)
// Combines RMSNorm normalization and FP8 quantization in a single kernel.
// This is optimal for the first layer where no residual connection exists.
void rms_norm_static_fp8_quant(
torch::Tensor& out, // [..., hidden_size], FP8 output
torch::Tensor& input, // [..., hidden_size], input tensor
torch::Tensor& weight, // [hidden_size], RMSNorm weight
torch::Tensor& scale, // [1], FP8 quantization scale
double epsilon); // RMSNorm epsilon
// Fused Add + RMSNorm + Static FP8 Quantization (with residual)
// Combines residual addition, RMSNorm, and FP8 quantization in a single kernel.
// The residual tensor is updated in-place with the sum of input and residual.
void fused_add_rms_norm_static_fp8_quant(
torch::Tensor& out, // [..., hidden_size], FP8 output
torch::Tensor& input, // [..., hidden_size], input tensor
torch::Tensor& residual, // [..., hidden_size], residual (updated in-place)
torch::Tensor& weight, // [hidden_size], RMSNorm weight
torch::Tensor& scale, // [1], FP8 quantization scale
double epsilon); // RMSNorm epsilon
// FP8 scaled matmul for W8A8 quantization using CUTLASS kernels
// Performs: c = (a @ b.T) with scales applied
torch::Tensor fp8_scaled_matmul(
const torch::Tensor& a,
const torch::Tensor& b,
const torch::Tensor& a_scale,
const torch::Tensor& b_scale,
torch::ScalarType output_dtype,
const std::optional<torch::Tensor>& bias = std::nullopt,
const std::optional<torch::Tensor>& output = std::nullopt);
std::pair<torch::Tensor, torch::Tensor> compute_topk_for_beam_search(
torch::Tensor combined_probs,
uint32_t batch_size,
uint32_t beam_size,
uint32_t top_k,
torch::Device device);
std::pair<torch::Tensor, torch::Tensor> compute_topk_general(
torch::Tensor input,
uint32_t batch_size,
uint32_t input_length,
uint32_t k,
torch::Device device);
torch::Tensor air_log_softmax_last_dim(const torch::Tensor& input,
const torch::Tensor& temperatures);
void fused_qk_norm_rope(
torch::Tensor& qkv, // Combined QKV tensor [num_tokens,
// (num_heads_q+num_heads_k+num_heads_v)*head_dim]
int64_t num_heads_q, // Number of query heads
int64_t num_heads_k, // Number of key heads
int64_t num_heads_v, // Number of value heads
int64_t head_dim, // Dimension per head
double eps, // Epsilon for RMS normalization
const torch::Tensor& q_weight, // RMSNorm weights for query [head_dim]
const torch::Tensor& k_weight, // RMSNorm weights for key [head_dim]
const torch::Tensor&
cos_sin_cache, // Cos/sin cache [max_position, rotary_dim]
bool interleaved, // Whether RoPE is applied in interleaved style
const torch::Tensor& position_ids // Position IDs for RoPE [num_tokens]
);
std::tuple<torch::Tensor, torch::Tensor> moe_fused_topk(
torch::Tensor& gating_output,
int64_t topk,
bool renormalize,
const std::optional<torch::Tensor>& correction_bias,
const std::string& scoring_func);
torch::Tensor random_sample(const torch::Tensor& probs);
torch::Tensor cutlass_fused_moe(
const torch::Tensor& input, // [num_tokens, hidden]
const torch::Tensor& token_selected_experts, // [num_tokens, top_k]
const torch::Tensor& token_final_scales, // [num_tokens, top_k]
const torch::Tensor&
fc1_expert_weights, // [num_experts, inter_dim, hidden]
const torch::Tensor&
fc2_expert_weights, // [num_experts, hidden, inter_dim]
torch::ScalarType output_dtype,
const std::vector<torch::Tensor>& quant_scales,
int32_t tp_size,
int32_t tp_rank,
int32_t ep_size,
int32_t ep_rank,
int32_t cluster_size,
int32_t cluster_rank,
const std::optional<torch::Tensor>& fc1_expert_biases = std::nullopt,
const std::optional<torch::Tensor>& fc2_expert_biases = std::nullopt,
const std::optional<torch::Tensor>& input_sf = std::nullopt,
const std::optional<torch::Tensor>& swiglu_alpha = std::nullopt,
const std::optional<torch::Tensor>& swiglu_beta = std::nullopt,
const std::optional<torch::Tensor>& swiglu_limit = std::nullopt,
const std::optional<torch::Tensor>& output = std::nullopt,
bool enable_alltoall = false,
bool use_deepseek_fp8_block_scale = false,
bool use_w4_group_scaling = false,
bool use_mxfp8_act_scaling = false,
bool min_latency_mode = false,
bool use_packed_weights = false,
int32_t tune_max_num_tokens = 8192,
ActivationType activation_type = ActivationType::SWIGLU);
// ---- moe_compute_index (moe_compute_index.cu) ----
// Fused routing index: bincount + argsort replacement.
// Returns {src_dst, dst_src, expert_sizes}.
std::tuple<torch::Tensor, torch::Tensor, torch::Tensor> moe_compute_index(
const torch::Tensor& expert_id,
int64_t num_experts);
// ---- moe_combine_result (moe_combine.cu) ----
// Fused combine: reorder + weighted sum in one pass.
torch::Tensor moe_combine_result(const torch::Tensor& gemm2,
const torch::Tensor& reduce_weight,
int64_t N,
int32_t topk);
} // namespace xllm::kernel::cuda

View File

@@ -0,0 +1,116 @@
/* Copyright 2025-2026 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
#if defined(USE_DCU)
#include <hip/amd_detail/amd_hip_bf16.h>
#include <hipcub/hipcub.hpp>
namespace cub = hipcub;
#else
#include <cub/cub.cuh>
#if CUB_VERSION >= 200800
#include <cuda/functional>
#endif
#endif
namespace xllm::kernel::cuda {
#if !defined(USE_DCU)
using BFloat16Type = __nv_bfloat16;
#define WARP_SIZE 32
#define XLLM_KERNEL_ATTR(MAX_THREADS)
#else
using BFloat16Type = hip_bfloat16;
#define WARP_SIZE 64
#define XLLM_KERNEL_ATTR(MAX_THREADS) __launch_bounds__(MAX_THREADS, 1)
#endif
#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))
template <typename T>
__device__ __forceinline__ T xllm_ldg(const T* ptr) {
#if defined(USE_DCU)
return *ptr;
#else
return __ldg(ptr);
#endif
}
// Define reduction operators based on CUB version.
#if defined(USE_DCU)
using MaxReduceOp = hipcub::Max;
using MinReduceOp = hipcub::Min;
#elif CUB_VERSION >= 200800
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);
#if defined(USE_DCU)
} else if constexpr (std::is_same_v<T, hip_bfloat16>) {
return __bfloat162float(reinterpret_cast<const __hip_bfloat16&>(x));
#else
} else if constexpr (std::is_same_v<T, __nv_bfloat16>) {
return __bfloat162float(x);
#endif
} 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

View File

@@ -0,0 +1,239 @@
/* Copyright 2025-2026 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
// clang-format off
#include <c10/util/Float8_e4m3fn.h>
#include <cmath>
#include <torch/types.h>
// clang-format on
namespace xllm {
namespace kernel {
namespace cuda {
// FP8 type max value definitions
template <typename T,
typename = std::enable_if_t<std::is_same_v<T, c10::Float8_e4m3fn> ||
std::is_same_v<T, int8_t>>>
struct quant_type_max {
static constexpr T val() { return std::numeric_limits<T>::max(); }
};
template <typename T>
__host__ __device__ static constexpr T quant_type_max_v =
quant_type_max<T>::val();
// Minimum scaling factor for quantization types
template <typename T,
typename = std::enable_if_t<std::is_same_v<T, c10::Float8_e4m3fn> ||
std::is_same_v<T, int8_t>>>
struct min_scaling_factor {
__device__ __host__ static inline float val() {
return 1.0f / (quant_type_max_v<T> * 512.0f);
}
};
template <>
struct min_scaling_factor<int8_t> {
__device__ __host__ static inline float val() {
return std::numeric_limits<float>::epsilon();
}
};
// Vectorization containers
template <typename scalar_t, size_t vec_size>
struct __align__(vec_size * sizeof(scalar_t)) vec_n_t {
scalar_t val[vec_size];
};
template <typename quant_type_t, size_t vec_size>
struct __align__(vec_size * sizeof(quant_type_t)) q8_n_t {
static_assert(std::is_same_v<quant_type_t, int8_t> ||
std::is_same_v<quant_type_t, c10::Float8_e4m3fn>);
quant_type_t val[vec_size];
};
// Atomic max for float
__device__ __forceinline__ float atomicMaxFloat(float* addr, float value) {
float old;
old = (value >= 0)
? __int_as_float(atomicMax((int*)addr, __float_as_int(value)))
: __uint_as_float(
atomicMin((unsigned int*)addr, __float_as_uint(value)));
return old;
}
// FP8 conversion functions
namespace fp8 {
#ifdef ENABLE_FP8
#include <cuda_fp8.h>
// float -> c10::Float8_e4m3fn conversion
template <typename Tout, typename Tin>
__inline__ __device__ Tout
vec_conversion(const Tin& x,
const __nv_fp8_interpretation_t fp8_type = __NV_E4M3) {
return x;
}
template <>
__inline__ __device__ c10::Float8_e4m3fn
vec_conversion<c10::Float8_e4m3fn, float>(
const float& a,
const __nv_fp8_interpretation_t fp8_type) {
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ < 800
return static_cast<c10::Float8_e4m3fn>(a);
#else
return c10::Float8_e4m3fn(__nv_cvt_float_to_fp8(a, __NV_SATFINITE, fp8_type),
c10::Float8_e4m3fn::from_bits());
#endif
}
#endif // ENABLE_FP8
} // namespace fp8
// Scaled FP8 conversion with saturation
template <bool is_scale_inverted, typename fp8_type>
__device__ __forceinline__ fp8_type scaled_fp8_conversion(float const val,
float const scale) {
float x = 0.0f;
if constexpr (is_scale_inverted) {
x = val * scale;
} else {
x = val / scale;
}
float r =
fmaxf(-quant_type_max_v<fp8_type>, fminf(x, quant_type_max_v<fp8_type>));
#ifdef ENABLE_FP8
// Use hardware cvt instruction for fp8 on nvidia
return fp8::vec_conversion<fp8_type, float>(r);
#else
return static_cast<fp8_type>(r);
#endif
}
// Vectorization utilities
template <int VEC_SIZE, typename InT, typename OutT, typename ScaOp>
struct DefaultVecOp {
ScaOp scalar_op;
__device__ __forceinline__ void operator()(
vec_n_t<OutT, VEC_SIZE>& dst,
const vec_n_t<InT, VEC_SIZE>& src) const {
#pragma unroll
for (int i = 0; i < VEC_SIZE; ++i) {
scalar_op(dst.val[i], src.val[i]);
}
}
};
template <int VEC_SIZE,
typename InT,
typename OutT,
typename VecOp,
typename ScaOp>
__device__ inline void vectorize_with_alignment(
const InT* in,
OutT* out,
int len,
int tid,
int stride,
VecOp&& vec_op, // vec_n_t<InT,16> -> vec_n_t<OutT,16>
ScaOp&& scalar_op) { // InT -> OutT
static_assert(VEC_SIZE > 0 && (VEC_SIZE & (VEC_SIZE - 1)) == 0,
"VEC_SIZE must be a positive power-of-two");
constexpr int WIDTH = VEC_SIZE * sizeof(InT);
uintptr_t addr = reinterpret_cast<uintptr_t>(in);
// Fast path when the whole region is already aligned
bool can_vec = ((addr & (WIDTH - 1)) == 0) && ((len & (VEC_SIZE - 1)) == 0);
if (can_vec) {
int num_vec = len / VEC_SIZE;
using vin_t = vec_n_t<InT, VEC_SIZE>;
using vout_t = vec_n_t<OutT, VEC_SIZE>;
auto* v_in = reinterpret_cast<const vin_t*>(in);
auto* v_out = reinterpret_cast<vout_t*>(out);
for (int i = tid; i < num_vec; i += stride) {
vout_t tmp;
vin_t src = v_in[i];
vec_op(tmp, src);
v_out[i] = tmp;
}
return;
}
int misalignment_offset = addr & (WIDTH - 1);
int alignment_bytes = WIDTH - misalignment_offset;
int prefix_elems = alignment_bytes & (WIDTH - 1);
prefix_elems /= sizeof(InT);
prefix_elems = min(prefix_elems, len);
// Prefix handling
for (int i = tid; i < prefix_elems; i += stride) {
scalar_op(out[i], in[i]);
}
in += prefix_elems;
out += prefix_elems;
len -= prefix_elems;
int num_vec = len / VEC_SIZE;
using vin_t = vec_n_t<InT, VEC_SIZE>;
using vout_t = vec_n_t<OutT, VEC_SIZE>;
auto* v_in = reinterpret_cast<const vin_t*>(in);
auto* v_out = reinterpret_cast<vout_t*>(out);
// Vectorized main part
for (int i = tid; i < num_vec; i += stride) {
vout_t tmp;
vin_t src = v_in[i];
vec_op(tmp, src);
v_out[i] = tmp;
}
// Tail handling
int tail_start = num_vec * VEC_SIZE;
for (int i = tid + tail_start; i < len; i += stride) {
scalar_op(out[i], in[i]);
}
}
template <int VEC_SIZE, typename InT, typename OutT, typename ScaOp>
__device__ __forceinline__ void vectorize_with_alignment(const InT* in,
OutT* out,
int len,
int tid,
int stride,
ScaOp&& scalar_op) {
using Vec = DefaultVecOp<VEC_SIZE, InT, OutT, std::decay_t<ScaOp>>;
vectorize_with_alignment<VEC_SIZE>(in,
out,
len,
tid,
stride,
Vec{scalar_op},
std::forward<ScaOp>(scalar_op));
}
} // namespace cuda
} // namespace kernel
} // namespace xllm

View File

@@ -0,0 +1,231 @@
/* Copyright 2025 The vLLM Authors and The xLLM Authors. All Rights Reserved.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://github.com/jd-opensource/xllm/blob/main/LICENSE
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <torch/all.h>
// ref to:
// https://github.com/vllm-project/vllm/blob/main/csrc/type_convert.cuh
/* Converter helpers for the conversion from torch types to HIP/CUDA types,
and the associated type conversions within HIP/CUDA. These helpers need
to be implemented for now because the relevant type conversion
operators/constructors are not consistently implemented by HIP/CUDA, so
a generic conversion via type casts cannot be implemented.
Each helper should have the member static constexpr bool `exists`:
If false, the optimized kernel is not used for the corresponding torch type.
If true, the helper should be fully defined as shown in the examples below.
*/
namespace xllm::kernel::cuda {
template <typename torch_type>
class _typeConvert {
public:
static constexpr bool exists = false;
};
template <>
class _typeConvert<float> {
public:
static constexpr bool exists = true;
using hip_type = float;
using packed_hip_type = float2;
using packed_hip_type4 = float4; // For 128-bit vectorization
__device__ static __forceinline__ float convert(hip_type x) { return x; }
__device__ static __forceinline__ float2 convert(packed_hip_type x) {
return x;
}
__device__ static __forceinline__ float4 convert(packed_hip_type4 x) {
return x;
}
};
#if defined(USE_DCU) || (defined(CUDA_VERSION) && (CUDA_VERSION >= 12000)) || \
defined(USE_MACA)
// CUDA < 12.0 runs into issues with packed type conversion
template <>
class _typeConvert<c10::Half> {
public:
static constexpr bool exists = true;
using hip_type = __half;
using packed_hip_type = __half2;
__device__ static __forceinline__ float convert(hip_type x) {
return __half2float(x);
}
__device__ static __forceinline__ float2 convert(packed_hip_type x) {
return __half22float2(x);
}
__device__ static __forceinline__ hip_type convert(float x) {
return __float2half_rn(x);
}
__device__ static __forceinline__ packed_hip_type convert(float2 x) {
return __float22half2_rn(x);
}
};
#endif // defined(USE_DCU) || CUDA_VERSION >= 12000
#if defined(USE_DCU)
template <>
class _typeConvert<c10::BFloat16> {
public:
static constexpr bool exists = true;
using hip_type = __hip_bfloat16;
using packed_hip_type = __hip_bfloat162;
__device__ static __forceinline__ float convert(hip_type x) {
return __bfloat162float(x);
}
__device__ static __forceinline__ float2 convert(packed_hip_type x) {
return __bfloat1622float2(x);
}
__device__ static __forceinline__ hip_type convert(float x) {
return __float2bfloat16(x);
}
__device__ static __forceinline__ packed_hip_type convert(float2 x) {
return __float22bfloat162_rn(x);
}
};
#elif defined(CUDA_VERSION) && (CUDA_VERSION >= 12000) && \
defined(__CUDA_ARCH__) && (__CUDA_ARCH__ >= 800) || \
defined(USE_MACA)
// CUDA_ARCH < 800 does not have BF16 support.
template <>
class _typeConvert<c10::BFloat16> {
public:
static constexpr bool exists = true;
using hip_type = __nv_bfloat16;
using packed_hip_type = __nv_bfloat162;
__device__ static __forceinline__ float convert(hip_type x) {
return __bfloat162float(x);
}
__device__ static __forceinline__ float2 convert(packed_hip_type x) {
return __bfloat1622float2(x);
}
__device__ static __forceinline__ hip_type convert(float x) {
return __float2bfloat16(x);
}
__device__ static __forceinline__ packed_hip_type convert(float2 x) {
return __float22bfloat162_rn(x);
}
};
#endif
/* Vector helper to generate vectorized and packed FP16/BF16 ops
for appropriate specializations of fused_add_rms_norm_kernel.
Only functions that are necessary in that kernel are implemented.
Alignment to 16 bytes is required to use 128-bit global memory ops.
*/
template <typename scalar_t, int width>
class alignas(16) _f16Vec {
public:
/* Not theoretically necessary that width is a power of 2 but should
almost always be the case for optimization purposes */
static_assert(width > 0 && (width & (width - 1)) == 0,
"Width is not a positive power of 2!");
using Converter = _typeConvert<scalar_t>;
using T1 = typename Converter::hip_type;
using T2 = typename Converter::packed_hip_type;
T1 data[width];
__device__ _f16Vec& operator+=(const _f16Vec<scalar_t, width>& other) {
if constexpr (width % 2 == 0) {
#pragma unroll
for (int i = 0; i < width; i += 2) {
if constexpr (std::is_same_v<T2, float2>) {
data[i] += other.data[i];
data[i + 1] += other.data[i + 1];
} else {
T2 temp{data[i], data[i + 1]};
temp += T2{other.data[i], other.data[i + 1]};
data[i] = temp.x;
data[i + 1] = temp.y;
}
}
} else {
#pragma unroll
for (int i = 0; i < width; ++i) data[i] += other.data[i];
}
return *this;
}
__device__ _f16Vec& operator*=(const _f16Vec<scalar_t, width>& other) {
if constexpr (width % 2 == 0) {
#pragma unroll
for (int i = 0; i < width; i += 2) {
if constexpr (std::is_same_v<T2, float2>) {
data[i] *= other.data[i];
data[i + 1] *= other.data[i + 1];
} else {
T2 temp{data[i], data[i + 1]};
temp *= T2{other.data[i], other.data[i + 1]};
data[i] = temp.x;
data[i + 1] = temp.y;
}
}
} else {
#pragma unroll
for (int i = 0; i < width; ++i) data[i] *= other.data[i];
}
return *this;
}
__device__ _f16Vec& operator*=(const float scale) {
if constexpr (width % 2 == 0) {
#pragma unroll
for (int i = 0; i < width; i += 2) {
float2 temp_f = Converter::convert(T2{data[i], data[i + 1]});
temp_f.x *= scale;
temp_f.y *= scale;
T2 temp = Converter::convert(temp_f);
data[i] = temp.x;
data[i + 1] = temp.y;
}
} else {
#pragma unroll
for (int i = 0; i < width; ++i) {
float temp = Converter::convert(data[i]) * scale;
data[i] = Converter::convert(temp);
}
}
return *this;
}
__device__ float sum_squares() const {
float result = 0.0f;
if constexpr (width % 2 == 0) {
#pragma unroll
for (int i = 0; i < width; i += 2) {
float2 z = Converter::convert(T2{data[i], data[i + 1]});
result += z.x * z.x + z.y * z.y;
}
} else {
#pragma unroll
for (int i = 0; i < width; ++i) {
float x = Converter::convert(data[i]);
result += x * x;
}
}
return result;
}
};
} // namespace xllm::kernel::cuda

View File

@@ -0,0 +1,163 @@
/* Copyright 2025-2026 The xLLM Authors.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://github.com/jd-opensource/xllm/blob/main/LICENSE
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include <ATen/DynamicLibrary.h>
#if defined(USE_DCU)
#include <c10/hip/HIPGuard.h>
#else
#include <c10/cuda/CUDAGuard.h>
#endif
#include <glog/logging.h>
#include <torch/torch.h>
#if !defined(USE_DCU)
#include <tvm/ffi/container/array.h>
#include <tvm/ffi/container/tensor.h>
#include <tvm/ffi/extra/c_env_api.h>
#include <tvm/ffi/extra/module.h>
#include <tvm/ffi/optional.h>
#endif
#include <string>
#include <tuple>
#include <type_traits>
#include <unordered_map>
#if defined(__CUDACC__) || defined(_NVHPC_CUDA) || defined(__HIPCC__)
#define HOST_DEVICE_INLINE __host__ __device__ __forceinline__
#define DEVICE_INLINE __device__ __forceinline__
#define HOST_INLINE __host__ __forceinline__
#else
#define HOST_DEVICE_INLINE inline
#define DEVICE_INLINE inline
#define HOST_INLINE inline
#endif
#if !defined(USE_DCU)
namespace ffi = tvm::ffi;
#endif
namespace xllm::kernel::cuda {
template <typename T>
HOST_DEVICE_INLINE constexpr std::enable_if_t<std::is_integral_v<T>, T>
ceil_div(T a, T b) {
return (a + b - 1) / b;
}
enum class ActivationType : int8_t {
GELU = 0,
RELU = 1,
SILU = 2,
SWIGLU = 3,
GEGLU = 4,
SWIGLU_BIAS = 5,
RELU2 = 6,
IDENTITY = 7,
INVALID_TYPE = 8
};
// torch tensor is only on cpu
torch::Tensor get_cache_buffer(const int32_t seq_len,
const torch::Device& device);
// NOLINTBEGIN(cppcoreguidelines-macro-usage)
#define DISPATCH_CASE_FLOATING_TYPES(...) \
AT_DISPATCH_CASE(at::ScalarType::Float, __VA_ARGS__) \
AT_DISPATCH_CASE(at::ScalarType::Half, __VA_ARGS__) \
AT_DISPATCH_CASE(at::ScalarType::BFloat16, __VA_ARGS__)
#define DISPATCH_FLOATING_TYPES(TYPE, NAME, ...) \
AT_DISPATCH_SWITCH(TYPE, NAME, DISPATCH_CASE_FLOATING_TYPES(__VA_ARGS__))
#define DISPATCH_CASE_HALF_TYPES(...) \
AT_DISPATCH_CASE(at::ScalarType::Half, __VA_ARGS__) \
AT_DISPATCH_CASE(at::ScalarType::BFloat16, __VA_ARGS__)
#define DISPATCH_HALF_TYPES(TYPE, NAME, ...) \
AT_DISPATCH_SWITCH(TYPE, NAME, DISPATCH_CASE_HALF_TYPES(__VA_ARGS__))
// NOLINTEND(cppcoreguidelines-macro-usage)
bool should_use_tensor_core(torch::ScalarType kv_cache_dtype,
int64_t num_attention_heads,
int64_t num_kv_heads);
bool support_pdl();
std::string path_to_uri_so_lib(const std::string& uri);
std::string determine_attention_backend(int64_t pos_encoding_mode,
bool use_fp16_qk_reduction,
bool use_custom_mask);
std::string get_batch_prefill_uri(const std::string& backend,
torch::ScalarType dtype_q,
torch::ScalarType dtype_kv,
torch::ScalarType dtype_o,
torch::ScalarType dtype_idx,
int64_t head_dim_qk,
int64_t head_dim_vo,
int64_t pos_encoding_mode,
bool use_sliding_window,
bool use_logits_soft_cap,
bool use_fp16_qk_reduction);
std::string get_batch_decode_uri(torch::ScalarType dtype_q,
torch::ScalarType dtype_kv,
torch::ScalarType dtype_o,
torch::ScalarType dtype_idx,
int64_t head_dim_qk,
int64_t head_dim_vo,
int64_t pos_encoding_mode,
bool use_sliding_window,
bool use_logits_soft_cap);
std::tuple<torch::Tensor, double> split_scale_param(const torch::Tensor& scale);
#if !defined(USE_DCU)
DLDataType to_dl_data_type(torch::ScalarType scalar_type);
// below are tvm-ffi related functions
ffi::Tensor to_ffi_tensor(const torch::Tensor& torch_tensor);
ffi::Optional<ffi::Tensor> to_ffi_optional_tensor(
const std::optional<torch::Tensor>& optional);
ffi::Array<ffi::Tensor> to_ffi_array_tensors(
const std::vector<torch::Tensor>& torch_tensors);
ffi::Optional<ffi::Array<ffi::Tensor>> to_ffi_optional_array_tensors(
const std::optional<std::vector<torch::Tensor>>& optional);
ffi::Module get_module(const std::string& uri);
ffi::Function get_function(const std::string& uri,
const std::string& func_name);
inline void bind_tvmffi_stream_to_current_torch_stream(
const torch::Device& device) {
const auto cur = c10::cuda::getCurrentCUDAStream(device.index());
// DLPack device type for CUDA is 2 (kDLCUDA).
void* original_stream = nullptr;
const int rc = TVMFFIEnvSetStream(
/*device_type=*/2,
/*device_id=*/device.index(),
reinterpret_cast<void*>(cur.stream()),
&original_stream);
if (rc != 0) {
LOG(WARNING) << "[tvmffi.stream] failed to set stream, rc=" << rc
<< " dev=" << device.index();
}
}
#endif // !defined(USE_DCU)
} // namespace xllm::kernel::cuda

View File

@@ -0,0 +1,600 @@
/* Copyright 2025 The vLLM Authors and 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.
==============================================================================*/
#include <c10/cuda/CUDAGuard.h>
#include <torch/cuda.h>
#include <cstdint>
#include <cub/cub.cuh>
#include "cuda_ops_api.h"
#include "device_utils.cuh"
#include "fp8_quant_utils.cuh"
#include "type_convert.cuh"
// ref to:
// https://github.com/vllm-project/vllm/blob/main/csrc/layernorm_kernels.cu
#if CUB_VERSION >= 200800
#include <cuda/std/functional>
using CubAddOp = ::cuda::std::plus<>;
using CubMaxOp = ::cuda::maximum<>;
#else // if CUB_VERSION < 200800
using CubAddOp = cub::Sum;
using CubMaxOp = cub::Max;
#endif // CUB_VERSION
namespace {
using namespace xllm::kernel::cuda;
template <typename scalar_t>
__global__ void XLLM_KERNEL_ATTR(1024)
rms_norm_kernel(scalar_t* __restrict__ out, // [..., hidden_size]
const scalar_t* __restrict__ input, // [..., hidden_size]
const int64_t input_stride,
const scalar_t* __restrict__ weight, // [hidden_size]
const float epsilon,
const int num_tokens,
const int hidden_size) {
__shared__ float s_variance;
float variance = 0.0f;
for (int idx = threadIdx.x; idx < hidden_size; idx += blockDim.x) {
const float x = static_cast<float>(input[blockIdx.x * input_stride + idx]);
variance += x * x;
}
using BlockReduce = cub::BlockReduce<float, 1024>;
__shared__ typename BlockReduce::TempStorage reduceStore;
variance = BlockReduce(reduceStore).Reduce(variance, CubAddOp{}, blockDim.x);
if (threadIdx.x == 0) {
s_variance = rsqrtf(variance / hidden_size + epsilon);
}
__syncthreads();
for (int idx = threadIdx.x; idx < hidden_size; idx += blockDim.x) {
float x = static_cast<float>(input[blockIdx.x * input_stride + idx]);
out[blockIdx.x * hidden_size + idx] =
(static_cast<scalar_t>(x * s_variance)) * weight[idx];
}
}
/* Function specialization in the case of FP16/BF16 tensors.
Additional optimizations we can make in this case are
packed and vectorized operations, which help with the
memory latency bottleneck. */
template <typename scalar_t, int width>
__global__ std::enable_if_t<(width > 0) && _typeConvert<scalar_t>::exists>
XLLM_KERNEL_ATTR(1024) fused_add_rms_norm_kernel(
scalar_t* __restrict__ input, // [..., hidden_size]
const int64_t input_stride,
scalar_t* __restrict__ residual, // [..., hidden_size]
const scalar_t* __restrict__ weight, // [hidden_size]
const float epsilon,
const int num_tokens,
const int hidden_size) {
// Sanity checks on our vector struct and type-punned pointer arithmetic
static_assert(std::is_pod_v<_f16Vec<scalar_t, width>>);
static_assert(sizeof(_f16Vec<scalar_t, width>) == sizeof(scalar_t) * width);
const int vec_hidden_size = hidden_size / width;
const int64_t vec_input_stride = input_stride / width;
__shared__ float s_variance;
float variance = 0.0f;
/* These and the argument pointers are all declared `restrict` as they are
not aliased in practice. Argument pointers should not be dereferenced
in this kernel as that would be undefined behavior */
auto* __restrict__ input_v =
reinterpret_cast<_f16Vec<scalar_t, width>*>(input);
auto* __restrict__ residual_v =
reinterpret_cast<_f16Vec<scalar_t, width>*>(residual);
auto* __restrict__ weight_v =
reinterpret_cast<const _f16Vec<scalar_t, width>*>(weight);
for (int idx = threadIdx.x; idx < vec_hidden_size; idx += blockDim.x) {
int id = blockIdx.x * vec_hidden_size + idx;
int64_t strided_id = blockIdx.x * vec_input_stride + idx;
_f16Vec<scalar_t, width> temp = input_v[strided_id];
temp += residual_v[id];
variance += temp.sum_squares();
residual_v[id] = temp;
}
using BlockReduce = cub::BlockReduce<float, 1024>;
__shared__ typename BlockReduce::TempStorage reduceStore;
variance = BlockReduce(reduceStore).Reduce(variance, CubAddOp{}, blockDim.x);
if (threadIdx.x == 0) {
s_variance = rsqrtf(variance / hidden_size + epsilon);
}
__syncthreads();
for (int idx = threadIdx.x; idx < vec_hidden_size; idx += blockDim.x) {
int id = blockIdx.x * vec_hidden_size + idx;
int64_t strided_id = blockIdx.x * vec_input_stride + idx;
_f16Vec<scalar_t, width> temp = residual_v[id];
temp *= s_variance;
temp *= weight_v[idx];
input_v[strided_id] = temp;
}
}
/* Generic fused_add_rms_norm_kernel
The width field is not used here but necessary for other specializations.
*/
template <typename scalar_t, int width>
__global__ std::enable_if_t<(width == 0) || !_typeConvert<scalar_t>::exists>
XLLM_KERNEL_ATTR(1024) fused_add_rms_norm_kernel(
scalar_t* __restrict__ input, // [..., hidden_size]
const int64_t input_stride,
scalar_t* __restrict__ residual, // [..., hidden_size]
const scalar_t* __restrict__ weight, // [hidden_size]
const float epsilon,
const int num_tokens,
const int hidden_size) {
__shared__ float s_variance;
float variance = 0.0f;
for (int idx = threadIdx.x; idx < hidden_size; idx += blockDim.x) {
scalar_t z = input[blockIdx.x * input_stride + idx];
z += residual[blockIdx.x * hidden_size + idx];
float x = static_cast<float>(z);
variance += x * x;
residual[blockIdx.x * hidden_size + idx] = z;
}
using BlockReduce = cub::BlockReduce<float, 1024>;
__shared__ typename BlockReduce::TempStorage reduceStore;
variance = BlockReduce(reduceStore).Reduce(variance, CubAddOp{}, blockDim.x);
if (threadIdx.x == 0) {
s_variance = rsqrtf(variance / hidden_size + epsilon);
}
__syncthreads();
for (int idx = threadIdx.x; idx < hidden_size; idx += blockDim.x) {
float x = static_cast<float>(residual[blockIdx.x * hidden_size + idx]);
input[blockIdx.x * input_stride + idx] =
(static_cast<scalar_t>(x * s_variance)) * weight[idx];
}
}
#define LAUNCH_FUSED_ADD_RMS_NORM(width) \
DISPATCH_FLOATING_TYPES( \
input.scalar_type(), "fused_add_rms_norm_kernel", [&] { \
fused_add_rms_norm_kernel<scalar_t, width> \
<<<grid, block, 0, stream>>>(input.data_ptr<scalar_t>(), \
input_stride, \
residual.data_ptr<scalar_t>(), \
weight.data_ptr<scalar_t>(), \
epsilon, \
num_tokens, \
hidden_size); \
});
// ============================================================================
// Fused RMSNorm + Static FP8 Quantization Kernels
// ============================================================================
// These kernels combine RMSNorm and FP8 quantization to reduce memory
// bandwidth by avoiding the intermediate write-back to global memory.
// Dispatch macro for FP8 types
#define DISPATCH_FP8_TYPES(TYPE, NAME, ...) \
[&] { \
const auto& the_type = TYPE; \
switch (the_type) { \
case at::ScalarType::Float8_e4m3fn: { \
using fp8_t = c10::Float8_e4m3fn; \
return __VA_ARGS__(); \
} \
default: \
AT_ERROR(#NAME, \
" not implemented for FP8 type '", \
toString(the_type), \
"'"); \
} \
}()
/**
* Fused RMSNorm + Static FP8 Quantization kernel (without residual)
* Combines RMSNorm and FP8 quantization in a single kernel to reduce
* memory bandwidth by avoiding intermediate write-back.
*
* @tparam scalar_t Input data type (float, half, bfloat16)
* @tparam fp8_type Output FP8 type (c10::Float8_e4m3fn)
* @param out Output FP8 tensor [num_tokens, hidden_size]
* @param input Input tensor [num_tokens, hidden_size]
* @param input_stride Stride of input tensor in the token dimension
* @param weight RMSNorm weight tensor [hidden_size]
* @param scale FP8 quantization scale (scalar)
* @param epsilon RMSNorm epsilon
* @param num_tokens Number of tokens
* @param hidden_size Hidden dimension size
*/
template <typename scalar_t, typename fp8_type>
__global__ void rms_norm_static_fp8_quant_kernel(
fp8_type* __restrict__ out, // [num_tokens, hidden_size]
const scalar_t* __restrict__ input, // [num_tokens, hidden_size]
const int64_t input_stride,
const scalar_t* __restrict__ weight, // [hidden_size]
const float* __restrict__ scale, // [1]
const float epsilon,
const int num_tokens,
const int hidden_size) {
__shared__ float s_variance;
float variance = 0.0f;
const scalar_t* input_row = input + blockIdx.x * input_stride;
// Step 1: Compute variance for RMSNorm
for (int idx = threadIdx.x; idx < hidden_size; idx += blockDim.x) {
const float x = static_cast<float>(input_row[idx]);
variance += x * x;
}
using BlockReduce = cub::BlockReduce<float, 1024>;
__shared__ typename BlockReduce::TempStorage reduceStore;
variance = BlockReduce(reduceStore).Reduce(variance, CubAddOp{}, blockDim.x);
if (threadIdx.x == 0) {
s_variance = rsqrtf(variance / hidden_size + epsilon);
}
__syncthreads();
// Step 2: Precompute scale inverse to avoid division
const float scale_inv = 1.0f / (*scale);
// Step 3: Fused RMSNorm + FP8 quantization
for (int idx = threadIdx.x; idx < hidden_size; idx += blockDim.x) {
float x = static_cast<float>(input_row[idx]);
float out_norm = (static_cast<scalar_t>(x * s_variance)) *
static_cast<float>(weight[idx]);
out[blockIdx.x * hidden_size + idx] =
xllm::kernel::cuda::scaled_fp8_conversion<true, fp8_type>(out_norm,
scale_inv);
}
}
/**
* Fused Add + RMSNorm + Static FP8 Quantization kernel (with residual)
* Optimized version with packed + vectorized operations for FP16/BF16.
*
* @tparam scalar_t Input data type (float, half, bfloat16)
* @tparam width Vector width for optimization (0, 8)
* @tparam fp8_type Output FP8 type (c10::Float8_e4m3fn)
*/
template <typename scalar_t, int width, typename fp8_type>
__global__ std::enable_if_t<(width > 0) && _typeConvert<scalar_t>::exists>
fused_add_rms_norm_static_fp8_quant_kernel(
fp8_type* __restrict__ out, // [num_tokens, hidden_size]
scalar_t* __restrict__ input, // [num_tokens, hidden_size]
const int64_t input_stride,
scalar_t* __restrict__ residual, // [num_tokens, hidden_size]
const scalar_t* __restrict__ weight, // [hidden_size]
const float* __restrict__ scale, // [1]
const float epsilon,
const int num_tokens,
const int hidden_size) {
static_assert(std::is_pod_v<_f16Vec<scalar_t, width>>);
static_assert(sizeof(_f16Vec<scalar_t, width>) == sizeof(scalar_t) * width);
const int vec_hidden_size = hidden_size / width;
const int64_t vec_input_stride = input_stride / width;
__shared__ float s_variance;
float variance = 0.0f;
auto* __restrict__ input_v =
reinterpret_cast<_f16Vec<scalar_t, width>*>(input);
auto* __restrict__ residual_v =
reinterpret_cast<_f16Vec<scalar_t, width>*>(residual);
auto* __restrict__ weight_v =
reinterpret_cast<const _f16Vec<scalar_t, width>*>(weight);
// Step 1: Fused add and compute variance
for (int idx = threadIdx.x; idx < vec_hidden_size; idx += blockDim.x) {
int id = blockIdx.x * vec_hidden_size + idx;
int64_t strided_id = blockIdx.x * vec_input_stride + idx;
_f16Vec<scalar_t, width> temp = input_v[strided_id];
temp += residual_v[id];
variance += temp.sum_squares();
residual_v[id] = temp; // Store updated residual
}
using BlockReduce = cub::BlockReduce<float, 1024>;
__shared__ typename BlockReduce::TempStorage reduceStore;
variance = BlockReduce(reduceStore).Reduce(variance, CubAddOp{}, blockDim.x);
if (threadIdx.x == 0) {
s_variance = rsqrtf(variance / hidden_size + epsilon);
}
__syncthreads();
// Step 2: Precompute scale inverse
const float scale_inv = 1.0f / (*scale);
// Step 3: Fused RMSNorm + FP8 quantization
for (int idx = threadIdx.x; idx < vec_hidden_size; idx += blockDim.x) {
int id = blockIdx.x * vec_hidden_size + idx;
_f16Vec<scalar_t, width> temp = residual_v[id];
temp *= s_variance;
temp *= weight_v[idx];
// Convert each element to FP8
#pragma unroll
for (int i = 0; i < width; ++i) {
float val = _typeConvert<scalar_t>::convert(temp.data[i]);
out[id * width + i] =
xllm::kernel::cuda::scaled_fp8_conversion<true, fp8_type>(val,
scale_inv);
}
}
}
/**
* Generic fused add + RMSNorm + FP8 quant kernel (fallback for unaligned data)
*/
template <typename scalar_t, int width, typename fp8_type>
__global__ std::enable_if_t<(width == 0) || !_typeConvert<scalar_t>::exists>
fused_add_rms_norm_static_fp8_quant_kernel(
fp8_type* __restrict__ out, // [num_tokens, hidden_size]
scalar_t* __restrict__ input, // [num_tokens, hidden_size]
const int64_t input_stride,
scalar_t* __restrict__ residual, // [num_tokens, hidden_size]
const scalar_t* __restrict__ weight, // [hidden_size]
const float* __restrict__ scale, // [1]
const float epsilon,
const int num_tokens,
const int hidden_size) {
__shared__ float s_variance;
float variance = 0.0f;
// Step 1: Fused add and compute variance
for (int idx = threadIdx.x; idx < hidden_size; idx += blockDim.x) {
scalar_t z = input[blockIdx.x * input_stride + idx];
z += residual[blockIdx.x * hidden_size + idx];
float x = static_cast<float>(z);
variance += x * x;
residual[blockIdx.x * hidden_size + idx] = z; // Store updated residual
}
using BlockReduce = cub::BlockReduce<float, 1024>;
__shared__ typename BlockReduce::TempStorage reduceStore;
variance = BlockReduce(reduceStore).Reduce(variance, CubAddOp{}, blockDim.x);
if (threadIdx.x == 0) {
s_variance = rsqrtf(variance / hidden_size + epsilon);
}
__syncthreads();
// Step 2: Precompute scale inverse
const float scale_inv = 1.0f / (*scale);
// Step 3: Fused RMSNorm + FP8 quantization
for (int idx = threadIdx.x; idx < hidden_size; idx += blockDim.x) {
float x = static_cast<float>(residual[blockIdx.x * hidden_size + idx]);
float out_norm = (static_cast<scalar_t>(x * s_variance)) *
static_cast<float>(weight[idx]);
out[blockIdx.x * hidden_size + idx] =
xllm::kernel::cuda::scaled_fp8_conversion<true, fp8_type>(out_norm,
scale_inv);
}
}
#define LAUNCH_FUSED_ADD_RMS_NORM_STATIC_FP8_QUANT(width) \
DISPATCH_FLOATING_TYPES( \
input.scalar_type(), "fused_add_rms_norm_static_fp8_quant", [&] { \
DISPATCH_FP8_TYPES( \
out.scalar_type(), "fused_add_rms_norm_static_fp8_quant", [&] { \
fused_add_rms_norm_static_fp8_quant_kernel<scalar_t, \
width, \
fp8_t> \
<<<grid, block, 0, stream>>>(out.data_ptr<fp8_t>(), \
input.data_ptr<scalar_t>(), \
input_stride, \
residual.data_ptr<scalar_t>(), \
weight.data_ptr<scalar_t>(), \
scale.data_ptr<float>(), \
epsilon, \
num_tokens, \
hidden_size); \
}); \
});
} // namespace
namespace xllm::kernel::cuda {
// flashinfer rmsnorm ops
// void rmsnorm(torch::Tensor output,
// torch::Tensor input,
// torch::Tensor weight,
// double eps) {
// FunctionFactory::get_instance().rmsnorm_func("norm").call(
// output, input, weight, eps, support_pdl());
// }
void rms_norm(torch::Tensor output, // [..., hidden_size]
torch::Tensor input, // [..., hidden_size]
torch::Tensor weight, // [hidden_size]
double eps) {
CHECK(output.is_contiguous());
CHECK(weight.is_contiguous());
// The kernel addresses tokens as `blockIdx.x * input_stride + idx`, which
// can only represent contiguous inputs or simple 2D strided rows. Flux q/k
// tensors reach this path as high-dimensional transposed views, so make that
// layout explicit before flattening tokens for the kernel.
if (input.dim() > 2 && !input.is_contiguous()) {
input = input.contiguous();
}
CHECK(input.stride(-1) == 1);
int hidden_size = input.size(-1);
int num_tokens = input.numel() / hidden_size;
int64_t input_stride = input.stride(-2);
dim3 grid(num_tokens);
dim3 block(std::min(hidden_size, 1024));
const at::cuda::OptionalCUDAGuard device_guard(device_of(input));
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
DISPATCH_FLOATING_TYPES(input.scalar_type(), "rms_norm_kernel", [&] {
rms_norm_kernel<scalar_t>
<<<grid, block, 0, stream>>>(output.data_ptr<scalar_t>(),
input.data_ptr<scalar_t>(),
input_stride,
weight.data_ptr<scalar_t>(),
eps,
num_tokens,
hidden_size);
});
}
void fused_add_rms_norm(torch::Tensor& input, // [..., hidden_size]
torch::Tensor& residual, // [..., hidden_size]
torch::Tensor& weight, // [hidden_size]
double epsilon) {
CHECK(weight.scalar_type() == input.scalar_type());
CHECK(input.scalar_type() == residual.scalar_type());
CHECK(residual.is_contiguous());
CHECK(weight.is_contiguous());
int hidden_size = input.size(-1);
int64_t input_stride = input.stride(-2);
int num_tokens = input.numel() / hidden_size;
dim3 grid(num_tokens);
/* This kernel is memory-latency bound in many scenarios.
When num_tokens is large, a smaller block size allows
for increased block occupancy on CUs and better latency
hiding on global mem ops. */
const int max_block_size = (num_tokens < 256) ? 1024 : 256;
dim3 block(std::min(hidden_size, max_block_size));
const at::cuda::OptionalCUDAGuard device_guard(device_of(input));
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
/*If the tensor types are FP16/BF16, try to use the optimized kernel
with packed + vectorized ops.
Max optimization is achieved with a width-8 vector of FP16/BF16s
since we can load at most 128 bits at once in a global memory op.
However, this requires each tensor's data to be aligned to 16
bytes.
*/
auto inp_ptr = reinterpret_cast<std::uintptr_t>(input.data_ptr());
auto res_ptr = reinterpret_cast<std::uintptr_t>(residual.data_ptr());
auto wt_ptr = reinterpret_cast<std::uintptr_t>(weight.data_ptr());
constexpr int kVectorWidth = 8;
constexpr int kReqAlignmentBytes =
kVectorWidth * 2; // kVectorWidth * sizeof(bfloat16 or float16) (float32
// falls back to non-vectorized version anyway)
bool ptrs_are_aligned = inp_ptr % kReqAlignmentBytes == 0 &&
res_ptr % kReqAlignmentBytes == 0 &&
wt_ptr % kReqAlignmentBytes == 0;
bool offsets_are_multiple_of_vector_width =
hidden_size % kVectorWidth == 0 && input_stride % kVectorWidth == 0;
if (ptrs_are_aligned && offsets_are_multiple_of_vector_width) {
LAUNCH_FUSED_ADD_RMS_NORM(8);
} else {
LAUNCH_FUSED_ADD_RMS_NORM(0);
}
}
// ============================================================================
// Fused RMSNorm + Static FP8 Quantization Host Functions
// ============================================================================
void rms_norm_static_fp8_quant(torch::Tensor& out, // [..., hidden_size], FP8
torch::Tensor& input, // [..., hidden_size]
torch::Tensor& weight, // [hidden_size]
torch::Tensor& scale, // [1]
double epsilon) {
CHECK(out.is_contiguous());
CHECK(input.stride(-1) == 1);
CHECK(weight.is_contiguous());
CHECK(scale.is_contiguous());
int hidden_size = input.size(-1);
int64_t input_stride = input.stride(-2);
int num_tokens = input.numel() / hidden_size;
// For large num_tokens, use smaller blocks to increase SM concurrency
const int max_block_size = (num_tokens < 256) ? 1024 : 256;
dim3 grid(num_tokens);
dim3 block(std::min(hidden_size, max_block_size));
const at::cuda::OptionalCUDAGuard device_guard(device_of(input));
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
DISPATCH_FLOATING_TYPES(
input.scalar_type(), "rms_norm_static_fp8_quant", [&] {
DISPATCH_FP8_TYPES(out.scalar_type(), "rms_norm_static_fp8_quant", [&] {
rms_norm_static_fp8_quant_kernel<scalar_t, fp8_t>
<<<grid, block, 0, stream>>>(out.data_ptr<fp8_t>(),
input.data_ptr<scalar_t>(),
input_stride,
weight.data_ptr<scalar_t>(),
scale.data_ptr<float>(),
epsilon,
num_tokens,
hidden_size);
});
});
}
void fused_add_rms_norm_static_fp8_quant(
torch::Tensor& out, // [..., hidden_size], FP8
torch::Tensor& input, // [..., hidden_size]
torch::Tensor& residual, // [..., hidden_size]
torch::Tensor& weight, // [hidden_size]
torch::Tensor& scale, // [1]
double epsilon) {
CHECK(out.is_contiguous());
CHECK(residual.is_contiguous());
CHECK(weight.is_contiguous());
CHECK(scale.is_contiguous());
CHECK(residual.scalar_type() == input.scalar_type());
CHECK(weight.scalar_type() == input.scalar_type());
int hidden_size = input.size(-1);
int64_t input_stride = input.stride(-2);
int num_tokens = input.numel() / hidden_size;
dim3 grid(num_tokens);
const int max_block_size = (num_tokens < 256) ? 1024 : 256;
dim3 block(std::min(hidden_size, max_block_size));
const at::cuda::OptionalCUDAGuard device_guard(device_of(input));
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
// Check alignment for vectorized kernel
auto inp_ptr = reinterpret_cast<std::uintptr_t>(input.data_ptr());
auto res_ptr = reinterpret_cast<std::uintptr_t>(residual.data_ptr());
auto wt_ptr = reinterpret_cast<std::uintptr_t>(weight.data_ptr());
constexpr int kVectorWidth = 8;
constexpr int kReqAlignmentBytes = kVectorWidth * 2;
bool ptrs_are_aligned = inp_ptr % kReqAlignmentBytes == 0 &&
res_ptr % kReqAlignmentBytes == 0 &&
wt_ptr % kReqAlignmentBytes == 0;
bool offsets_are_multiple_of_vector_width =
hidden_size % kVectorWidth == 0 && input_stride % kVectorWidth == 0;
if (ptrs_are_aligned && offsets_are_multiple_of_vector_width) {
LAUNCH_FUSED_ADD_RMS_NORM_STATIC_FP8_QUANT(8);
} else {
LAUNCH_FUSED_ADD_RMS_NORM_STATIC_FP8_QUANT(0);
}
}
} // namespace xllm::kernel::cuda

View File

@@ -0,0 +1,101 @@
/* Copyright 2025-2026 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.
==============================================================================*/
#include <c10/cuda/CUDAStream.h>
#include "cuda_ops_api.h"
#include "device_utils.cuh"
namespace xllm::kernel::cuda {
template <typename T>
__global__ void XLLM_KERNEL_ATTR(1024) reshape_paged_cache_kernel(
const int* __restrict__ slot_ids, // [n_tokens]
const T* __restrict__ keys, // [n_tokens, n_heads, head_dim]
const T* __restrict__ values, // [n_tokens, n_heads, head_dim]
T* __restrict__ key_cache,
T* __restrict__ value_cache,
int64_t k_stride,
int64_t v_stride,
int64_t n_kv_heads,
int64_t head_dim,
int64_t block_size) {
// block/token index
const int64_t bid = blockIdx.x;
// which slot to write to
const int64_t slot_id = slot_ids[bid];
if (slot_id < 0) {
return;
}
// block index
const int64_t block_idx = slot_id / block_size;
// offset within block
const int64_t block_offset = slot_id % block_size;
// base index for the block in cache
const int64_t block_base_idx = block_idx * block_size * n_kv_heads * head_dim;
// copy value one by one for the token
for (int64_t i = threadIdx.x; i < n_kv_heads * head_dim; i += blockDim.x) {
const int64_t k_src_idx = bid * k_stride + i;
const int64_t v_src_idx = bid * v_stride + i;
// cache: [n_blocks, block_size, n_heads, head_dim]
const int64_t head_base_idx =
block_base_idx + block_offset * n_kv_heads * head_dim;
// which head to write to
const int head_idx = i / head_dim;
// which dim within head to write to
const int head_offset = i % head_dim;
const int64_t dst_idx = head_base_idx + head_idx * head_dim + head_offset;
key_cache[dst_idx] = keys[k_src_idx];
value_cache[dst_idx] = values[v_src_idx];
}
}
void reshape_paged_cache(
torch::Tensor slot_ids, // [n_tokens]
torch::Tensor keys, // [n_tokens, n_kv_heads, head_dim]
torch::Tensor values, // [n_tokens, n_kv_heads, head_dim]
torch::Tensor key_cache, // [n_blocks, block_size, n_heads, head_dim]
torch::Tensor value_cache) {
// keys and values should be continuous at n_kv_heads and head_dim dims
CHECK(keys.stride(-1) == 1 && keys.stride(-2) == keys.size(-1));
CHECK(values.stride(-1) == 1 && values.stride(-2) == values.size(-1));
const int64_t n_tokens = keys.size(-3);
const int64_t n_kv_heads = keys.size(-2);
const int64_t head_dim = keys.size(-1);
const int64_t block_size = key_cache.size(-3);
// it is possible that keys and values have different strides
const int64_t k_stride = keys.stride(-3);
const int64_t v_stride = values.stride(-3);
const int64_t n = n_kv_heads * head_dim;
dim3 grid(n_tokens);
dim3 block(std::min<int>(n, 1024));
DISPATCH_FLOATING_TYPES(
keys.scalar_type(), "reshape_paged_cache_kernel", [&] {
reshape_paged_cache_kernel<scalar_t>
<<<grid, block, 0, c10::cuda::getCurrentCUDAStream()>>>(
slot_ids.data_ptr<int>(),
keys.data_ptr<scalar_t>(),
values.data_ptr<scalar_t>(),
key_cache.data_ptr<scalar_t>(),
value_cache.data_ptr<scalar_t>(),
k_stride,
v_stride,
n_kv_heads,
head_dim,
block_size);
});
}
} // namespace xllm::kernel::cuda

View File

@@ -0,0 +1,258 @@
/* Copyright 2025 The vLLM Authors and 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.
==============================================================================*/
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <torch/all.h>
#include "cuda_ops_api.h"
#include "device_utils.cuh"
// ref to:
// https://github.com/vllm-project/vllm/blob/main/csrc/pos_encoding_kernels.cu
namespace {
template <typename scalar_t, bool IS_NEOX>
inline __device__ void apply_token_rotary_embedding(
scalar_t* __restrict__ arr,
const scalar_t* __restrict__ cos_ptr,
const scalar_t* __restrict__ sin_ptr,
int rot_offset,
int embed_dim) {
int x_index, y_index;
scalar_t cos, sin;
if (IS_NEOX) {
// GPT-NeoX style rotary embedding.
x_index = rot_offset;
y_index = embed_dim + rot_offset;
cos = *(cos_ptr + x_index);
sin = *(sin_ptr + x_index);
} else {
// GPT-J style rotary embedding.
x_index = 2 * rot_offset;
y_index = 2 * rot_offset + 1;
cos = *(cos_ptr + x_index / 2);
sin = *(sin_ptr + x_index / 2);
}
const scalar_t x = arr[x_index];
const scalar_t y = arr[y_index];
arr[x_index] = x * cos - y * sin;
arr[y_index] = y * cos + x * sin;
}
template <typename scalar_t, bool IS_NEOX>
inline __device__ void apply_rotary_embedding(
scalar_t* __restrict__ query, // [batch_size, seq_len, num_heads,
// head_size] or [num_tokens, num_heads,
// head_size]
scalar_t* __restrict__ key, // nullptr or
// [batch_size, seq_len, num_kv_heads,
// head_size] or [num_tokens, num_kv_heads,
// head_size]
const scalar_t* cache_ptr,
const int head_size,
const int num_heads,
const int num_kv_heads,
const int rot_dim,
const int token_idx,
const int64_t query_stride,
const int64_t key_stride,
const int64_t head_stride) {
const int embed_dim = rot_dim / 2;
const scalar_t* cos_ptr = cache_ptr;
const scalar_t* sin_ptr = cache_ptr + embed_dim;
const int nq = num_heads * embed_dim;
for (int i = threadIdx.x; i < nq; i += blockDim.x) {
const int head_idx = i / embed_dim;
const int64_t token_head =
token_idx * query_stride + head_idx * head_stride;
const int rot_offset = i % embed_dim;
apply_token_rotary_embedding<scalar_t, IS_NEOX>(
query + token_head, cos_ptr, sin_ptr, rot_offset, embed_dim);
}
if (key != nullptr) {
const int nk = num_kv_heads * embed_dim;
for (int i = threadIdx.x; i < nk; i += blockDim.x) {
const int head_idx = i / embed_dim;
const int64_t token_head =
token_idx * key_stride + head_idx * head_stride;
const int rot_offset = i % embed_dim;
apply_token_rotary_embedding<scalar_t, IS_NEOX>(
key + token_head, cos_ptr, sin_ptr, rot_offset, embed_dim);
}
}
}
template <typename scalar_t, bool IS_NEOX>
__global__ void XLLM_KERNEL_ATTR(512) rotary_embedding_kernel(
const int64_t* __restrict__ positions, // [batch_size, seq_len] or
// [num_tokens]
scalar_t* __restrict__ query, // [batch_size, seq_len, num_heads,
// head_size] or [num_tokens, num_heads,
// head_size]
scalar_t* __restrict__ key, // nullptr or
// [batch_size, seq_len, num_kv_heads,
// head_size] or [num_tokens, num_kv_heads,
// head_size]
const scalar_t* __restrict__ cos_sin_cache, // [max_position, 2,
// rot_dim // 2]
const int rot_dim,
const int64_t query_stride,
const int64_t key_stride,
const int64_t head_stride,
const int num_heads,
const int num_kv_heads,
const int head_size) {
// Each thread block is responsible for one token.
const int token_idx = blockIdx.x;
int64_t pos = positions[token_idx];
const scalar_t* cache_ptr = cos_sin_cache + pos * rot_dim;
apply_rotary_embedding<scalar_t, IS_NEOX>(query,
key,
cache_ptr,
head_size,
num_heads,
num_kv_heads,
rot_dim,
token_idx,
query_stride,
key_stride,
head_stride);
}
} // namespace
namespace xllm::kernel::cuda {
// flashinfer rope ops
// void apply_rope_pos_ids_cos_sin_cache(torch::Tensor q,
// torch::Tensor k,
// torch::Tensor cos_sin_cache,
// torch::Tensor pos_ids,
// bool interleave) {
// const int64_t head_dim = cos_sin_cache.size(-1) / 2;
// q = q.view({q.size(0), -1, head_dim});
// k = k.view({k.size(0), -1, head_dim});
// FunctionFactory::get_instance().rope_func("rope").call(
// q, k, q, k, cos_sin_cache, pos_ids, interleave);
// }
void rotary_embedding(
torch::Tensor& positions, // [batch_size, seq_len] or [num_tokens]
torch::Tensor& query, // [batch_size, seq_len, num_heads * head_size] or
// [num_tokens, num_heads * head_size] or
// [batch_size, seq_len, num_heads, head_size] or
// [num_tokens, num_heads, head_size]
std::optional<torch::Tensor> key,
// null or
// [batch_size, seq_len, num_kv_heads * head_size] or
// [num_tokens, num_kv_heads * head_size] or
// [batch_size, seq_len, num_heads, head_size] or
// [num_tokens, num_heads, head_size]
// int64_t head_size,
torch::Tensor& cos_sin_cache, // [max_position, rot_dim]
bool is_neox) {
// num_tokens = batch_size * seq_len
const int positions_ndim = positions.dim();
const int query_ndim = query.dim();
// For partial rotary models, e.g. MiniMax-M2 with head_dim=128 and
// rotary_dim=64, the cache width is the rotary dimension rather than the
// physical per-head stride. When query is already shaped as
// [*, num_heads, head_size], infer the real head_size from query itself.
int64_t head_size = (query_ndim == positions_ndim + 2)
? query.size(-1)
: cos_sin_cache.size(-1);
int64_t num_tokens = positions.numel();
// Make sure num_tokens dim is consistent across positions, query, and key
CHECK(positions_ndim == 1 || positions_ndim == 2)
<< "positions must have shape [num_tokens] or [batch_size, seq_len]";
if (positions_ndim == 1) {
CHECK(query.size(0) == positions.size(0) &&
(!key.has_value() || key->size(0) == positions.size(0)))
<< "query, key and positions must have the same number of tokens";
}
if (positions_ndim == 2) {
CHECK(query.size(0) == positions.size(0) &&
(!key.has_value() || key->size(0) == positions.size(0)) &&
query.size(1) == positions.size(1) &&
(!key.has_value() || key->size(1) == positions.size(1)))
<< "query, key and positions must have the same batch_size and seq_len";
}
// Make sure head_size is valid for query and key
// hidden_size = num_heads * head_size
int query_hidden_size = query.numel() / num_tokens;
int key_hidden_size = key.has_value() ? key->numel() / num_tokens : 0;
CHECK(query_hidden_size % head_size == 0);
CHECK(key_hidden_size % head_size == 0);
// Make sure query and key have consistent number of heads
int num_heads = query_hidden_size / head_size;
int num_kv_heads = key.has_value() ? key_hidden_size / head_size : num_heads;
CHECK(num_heads % num_kv_heads == 0);
int rot_dim = cos_sin_cache.size(1);
int seq_dim_idx = positions_ndim - 1;
int64_t query_stride = query.stride(seq_dim_idx);
int64_t key_stride = key.has_value() ? key->stride(seq_dim_idx) : 0;
// Determine head stride: for [*, heads, head_size] use stride of last dim;
// for flat [*, heads*head_size], heads blocks are contiguous of size
// head_size
int64_t head_stride =
(query_ndim == positions_ndim + 2) ? query.stride(-2) : head_size;
dim3 grid(num_tokens);
dim3 block(std::min<int64_t>(num_heads * rot_dim / 2, 512));
const at::cuda::OptionalCUDAGuard device_guard(device_of(query));
const cudaStream_t stream = at::cuda::getCurrentCUDAStream();
DISPATCH_FLOATING_TYPES(
query.scalar_type(), "apply_rope_pos_ids_cos_sin_cache", [&] {
if (is_neox) {
rotary_embedding_kernel<scalar_t, true><<<grid, block, 0, stream>>>(
positions.data_ptr<int64_t>(),
query.data_ptr<scalar_t>(),
key.has_value() ? key->data_ptr<scalar_t>() : nullptr,
cos_sin_cache.data_ptr<scalar_t>(),
rot_dim,
query_stride,
key_stride,
head_stride,
num_heads,
num_kv_heads,
head_size);
} else {
rotary_embedding_kernel<scalar_t, false><<<grid, block, 0, stream>>>(
positions.data_ptr<int64_t>(),
query.data_ptr<scalar_t>(),
key.has_value() ? key->data_ptr<scalar_t>() : nullptr,
cos_sin_cache.data_ptr<scalar_t>(),
rot_dim,
query_stride,
key_stride,
head_stride,
num_heads,
num_kv_heads,
head_size);
}
});
}
} // namespace xllm::kernel::cuda

View File

@@ -0,0 +1,32 @@
/* Copyright 2025-2026 The xLLM Authors.
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.
==============================================================================*/
#include "ilu_ops_api.h"
using namespace ixformer;
namespace xllm::kernel::ilu {
void act_and_mul(torch::Tensor out,
torch::Tensor input,
const std::string& act_mode) {
if (act_mode == "silu") {
infer::silu_and_mul(input, out);
} else {
LOG(FATAL) << "Unsupported act mode: " << act_mode
<< ", only support silu, gelu, gelu_tanh";
}
}
} // namespace xllm::kernel::ilu

View File

@@ -0,0 +1,163 @@
/* Copyright 2025-2026 The xLLM Authors.
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.
==============================================================================*/
#include "ilu_ops_api.h"
#include "ixinfer.h"
#include "utils.h"
using namespace ixformer;
namespace xllm::kernel::ilu {
void reshape_paged_cache(torch::Tensor& key,
std::optional<torch::Tensor>& value,
torch::Tensor& key_cache,
std::optional<torch::Tensor>& value_cache,
torch::Tensor& slot_mapping) {
auto value_ = value.value_or(torch::Tensor());
auto value_cache_ = value_cache.value_or(torch::Tensor());
int64_t key_token_stride = key.stride(0);
int64_t value_token_stride = 0;
if (value_.defined()) {
value_token_stride = value_.stride(0);
}
slot_mapping = slot_mapping.to(at::kLong);
infer::xllm_reshape_and_cache(key,
value_,
key_cache,
value_cache_,
slot_mapping,
key_token_stride,
value_token_stride);
}
void batch_prefill(torch::Tensor& query,
const torch::Tensor& key,
const std::optional<torch::Tensor>& value,
torch::Tensor& output,
std::optional<torch::Tensor>& output_lse,
const std::optional<torch::Tensor>& q_cu_seq_lens,
const std::optional<torch::Tensor>& kv_cu_seq_lens,
const std::optional<torch::Tensor>& alibi_slope,
const std::optional<torch::Tensor>& attn_bias,
const std::optional<torch::Tensor>& q_quant_scale,
const std::optional<torch::Tensor>& k_quant_scale,
const std::optional<torch::Tensor>& v_quant_scale,
const torch::Tensor& block_tables,
int64_t max_query_len,
int64_t max_seq_len,
float scale,
bool is_causal,
int64_t window_size_left,
int64_t window_size_right,
const std::string& compute_dtype,
bool return_lse) {
double softcap = 0.0;
bool sqrt_alibi = false;
auto q_cu_seq_lens_ = q_cu_seq_lens.value_or(torch::Tensor());
auto kv_cu_seq_lens_ = kv_cu_seq_lens.value_or(torch::Tensor());
auto q_quant_scale_ = q_quant_scale.value_or(torch::Tensor());
auto k_quant_scale_ = k_quant_scale.value_or(torch::Tensor());
auto v_quant_scale_ = v_quant_scale.value_or(torch::Tensor());
auto block_tables_ = block_tables;
auto key_ = key;
auto value_ = value.value();
infer::ixinfer_flash_attn_unpad_with_block_tables(query,
key_,
value_,
output,
block_tables_,
q_cu_seq_lens_,
kv_cu_seq_lens_,
max_query_len,
max_seq_len,
is_causal,
window_size_left,
window_size_right,
static_cast<double>(scale),
softcap,
sqrt_alibi,
alibi_slope,
c10::nullopt,
output_lse);
}
void batch_decode(torch::Tensor& query,
const torch::Tensor& k_cache,
torch::Tensor& output,
const torch::Tensor& block_table,
const torch::Tensor& seq_lens,
const std::optional<torch::Tensor>& v_cache,
std::optional<torch::Tensor>& output_lse,
const std::optional<torch::Tensor>& q_quant_scale,
const std::optional<torch::Tensor>& k_cache_quant_scale,
const std::optional<torch::Tensor>& v_cache_quant_scale,
const std::optional<torch::Tensor>& out_quant_scale,
const std::optional<torch::Tensor>& alibi_slope,
const std::optional<torch::Tensor>& mask,
const std::string& compute_dtype,
int64_t max_seq_len,
int64_t window_size_left,
int64_t window_size_right,
float scale,
bool return_lse,
bool is_causal,
int64_t kv_cache_quant_bit_size) {
if (query.dim() == 4) {
query =
query
.view({query.size(0) * query.size(1), query.size(2), query.size(3)})
.contiguous();
}
if (output.dim() == 4) {
output = output
.view({output.size(0) * output.size(1),
output.size(2),
output.size(3)})
.contiguous();
;
}
auto v_cache_ = v_cache.value_or(torch::Tensor());
int64_t num_kv_heads = k_cache.size(1);
int64_t page_block_size = k_cache.size(2);
double softcap = 0.0;
bool enable_cuda_graph = false;
bool use_sqrt_alibi = false;
auto block_table_ = block_table;
auto k_cache_ = k_cache;
auto seq_lens_ = seq_lens;
infer::xllm_paged_attention(output,
query,
k_cache_,
v_cache_,
num_kv_heads,
scale,
block_table_,
seq_lens_,
page_block_size,
max_seq_len,
alibi_slope,
is_causal,
(int32_t)window_size_left,
(int32_t)window_size_right,
softcap,
enable_cuda_graph,
use_sqrt_alibi,
c10::nullopt);
}
} // namespace xllm::kernel::ilu

View File

@@ -0,0 +1,99 @@
/* Copyright 2025-2026 The xLLM Authors.
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.
==============================================================================*/
#include <glog/logging.h>
#include "ilu_ops_api.h"
namespace xllm::kernel::ilu {
std::tuple<torch::Tensor, torch::Tensor> moe_active_topk(
const torch::Tensor& input,
int64_t topk,
int64_t num_expert_group,
int64_t topk_group,
bool normalize,
const std::optional<torch::Tensor>& mask,
const std::string& normed_by,
const std::string& scoring_func,
double route_scale,
const std::optional<torch::Tensor>& e_score_correction_bias) {
torch::Tensor input_ = input.to(torch::kFloat32);
auto reduce_weight =
torch::empty({input.size(0), topk},
torch::dtype(torch::kFloat).device(input.device()));
auto topk_indices =
torch::empty({input.size(0), topk},
torch::dtype(torch::kInt32).device(input.device()));
auto token_expert_indices =
torch::empty({input.size(0), topk},
torch::dtype(torch::kInt32).device(input.device()));
infer::topk_softmax(
reduce_weight, topk_indices, token_expert_indices, input_, false);
auto tt = reduce_weight.sum(-1);
if (normalize) {
reduce_weight = reduce_weight / reduce_weight.sum(-1).unsqueeze(-1);
}
return std::make_tuple(reduce_weight, topk_indices);
}
std::vector<torch::Tensor> moe_gen_idx(torch::Tensor& expert_id,
int64_t expert_num) {
auto src_dst = expert_id.new_empty({expert_id.numel()});
auto dst_src = torch::empty_like(src_dst);
auto expert_sizes_gpu = expert_id.new_empty({expert_num});
auto expert_sizes_gpu_cumsum = expert_id.new_zeros({expert_id.numel() + 1});
infer::moe_compute_token_index_api(expert_id,
src_dst,
dst_src,
expert_sizes_gpu,
/*expert_mask=*/std::nullopt,
/*expert_sizes_cpu*/ std::nullopt,
/*expert_sizes_gpu*/ std::nullopt,
0,
expert_num,
expert_num);
expert_sizes_gpu_cumsum = expert_sizes_gpu.cumsum(-1);
return {src_dst, dst_src, expert_sizes_gpu, expert_sizes_gpu_cumsum};
}
torch::Tensor moe_expand_input(const torch::Tensor& input,
const torch::Tensor& gather_index,
const torch::Tensor& combine_idx,
int64_t topk) {
int64_t dst_tokens = input.size(0) * topk;
auto output = input.new_empty({dst_tokens, input.size(1)});
infer::moe_expand_input(
output, input, combine_idx, gather_index, dst_tokens, topk);
return output;
}
torch::Tensor moe_combine_result(torch::Tensor& input, torch::Tensor& weight) {
input = input.view({-1, weight.size(1), input.size(1)});
auto output = input.new_empty({input.size(0), input.size(2)});
infer::moe_output_reduce_sum(output,
input,
weight,
/*mask=*/std::nullopt,
/*extra_residual*/ std::nullopt,
/*scaling_factor=*/1.0);
return output;
}
} // namespace xllm::kernel::ilu

View File

@@ -0,0 +1,39 @@
/* Copyright 2025-2026 The xLLM Authors.
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.
==============================================================================*/
#include "ilu_ops_api.h"
namespace xllm::kernel::ilu {
torch::Tensor group_gemm(torch::Tensor& input,
torch::Tensor& weight,
torch::Tensor& tokens_per_experts,
const std::optional<torch::Tensor>& dst_to_src,
torch::Tensor& output) {
infer::moe_w16a16_group_gemm(
output,
input,
weight,
tokens_per_experts,
dst_to_src,
/*bias=*/std::nullopt,
/*format=*/"TN",
/*persistent=*/0,
/*output_n=*/tokens_per_experts.sum().item<int64_t>());
return output;
}
} // namespace xllm::kernel::ilu

View File

@@ -0,0 +1,153 @@
/* Copyright 2025-2026 The xLLM Authors.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
https://github.com/jd-opensource/xllm/blob/main/LICENSE
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#pragma once
#include <ATen/DynamicLibrary.h>
#include <ATen/core/dispatch/Dispatcher.h>
#include <cuda_runtime.h>
#include <glog/logging.h>
#include <torch/all.h>
#include <optional>
#include "ATen/Tensor.h"
#include "ATen/cuda/CUDAEvent.h"
#include "c10/core/Device.h"
#include "c10/core/DeviceGuard.h"
#include "c10/core/GradMode.h"
#include "c10/core/InferenceMode.h"
#include "c10/core/MemoryFormat.h"
#include "c10/core/ScalarType.h"
#include "c10/core/TensorOptions.h"
#include "c10/cuda/CUDAFunctions.h"
#include "c10/cuda/CUDAGuard.h"
#include "c10/cuda/CUDAStream.h"
#include "ixformer.h"
#include "kernels/kernels.h"
// #include "utils.h"
using namespace ixformer;
namespace xllm::kernel::ilu {
void apply_rope_pos_ids_cos_sin_cache(torch::Tensor& query,
torch::Tensor& key,
torch::Tensor& cos_sin_cache,
torch::Tensor& positions,
bool interleave);
// act_mode only support silu, gelu, gelu_tanh
void act_and_mul(torch::Tensor out,
torch::Tensor input,
const std::string& act_mode);
void reshape_paged_cache(
torch::Tensor& key, // (num_tokens, num_heads, head_size)
std::optional<torch::Tensor>& value, // (num_tokens, num_heads, head_size)
torch::Tensor& key_cache, // (num_blocks, num_heads, block_size, head_size)
std::optional<torch::Tensor>&
value_cache, // (num_blocks, num_heads, block_size, head_size)
torch::Tensor& slot_mapping); //(num_tokens)
void batch_prefill(torch::Tensor& query,
const torch::Tensor& key,
const std::optional<torch::Tensor>& value,
torch::Tensor& output,
std::optional<torch::Tensor>& output_lse,
const std::optional<torch::Tensor>& q_cu_seq_lens,
const std::optional<torch::Tensor>& kv_cu_seq_lens,
const std::optional<torch::Tensor>& alibi_slope,
const std::optional<torch::Tensor>& attn_bias,
const std::optional<torch::Tensor>& q_quant_scale,
const std::optional<torch::Tensor>& k_quant_scale,
const std::optional<torch::Tensor>& v_quant_scale,
const torch::Tensor& block_tables,
int64_t max_query_len,
int64_t max_seq_len,
float scale,
bool is_causal,
int64_t window_size_left,
int64_t window_size_right,
const std::string& compute_dtype,
bool return_lse);
void batch_decode(torch::Tensor& query,
const torch::Tensor& k_cache,
torch::Tensor& output,
const torch::Tensor& block_table,
const torch::Tensor& seq_lens,
const std::optional<torch::Tensor>& v_cache,
std::optional<torch::Tensor>& output_lse,
const std::optional<torch::Tensor>& q_quant_scale,
const std::optional<torch::Tensor>& k_cache_quant_scale,
const std::optional<torch::Tensor>& v_cache_quant_scale,
const std::optional<torch::Tensor>& out_quant_scale,
const std::optional<torch::Tensor>& alibi_slope,
const std::optional<torch::Tensor>& mask,
const std::string& compute_dtype,
int64_t max_seq_len,
int64_t window_size_left,
int64_t window_size_right,
float scale,
bool return_lse,
bool is_causal,
int64_t kv_cache_quant_bit_size);
void residual_layer_norm(torch::Tensor& input,
torch::Tensor& output,
std::optional<torch::Tensor>& residual,
torch::Tensor& weight,
std::optional<torch::Tensor>& bias,
std::optional<torch::Tensor>& residual_out,
double eps);
void rms_norm(torch::Tensor& output,
torch::Tensor& input,
torch::Tensor& weight,
double eps);
torch::Tensor matmul(torch::Tensor a,
torch::Tensor b,
std::optional<torch::Tensor> bias);
std::tuple<torch::Tensor, torch::Tensor> moe_active_topk(
const torch::Tensor& input,
int64_t topk,
int64_t num_expert_group,
int64_t topk_group,
bool normalize,
const std::optional<torch::Tensor>& mask,
const std::string& normed_by,
const std::string& scoring_func,
double route_scale,
const std::optional<torch::Tensor>& e_score_correction_bias);
std::vector<torch::Tensor> moe_gen_idx(torch::Tensor& expert_id,
int64_t expert_num);
torch::Tensor moe_expand_input(const torch::Tensor& input,
const torch::Tensor& gather_index,
const torch::Tensor& combine_idx,
int64_t topk);
torch::Tensor group_gemm(torch::Tensor& input,
torch::Tensor& weight,
torch::Tensor& tokens_per_experts,
const std::optional<torch::Tensor>& dst_to_src,
torch::Tensor& output);
torch::Tensor moe_combine_result(torch::Tensor& input, torch::Tensor& weight);
} // namespace xllm::kernel::ilu

View File

@@ -0,0 +1,147 @@
/* Copyright 2025-2026 The xLLM Authors.
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.
==============================================================================*/
#include <torch/all.h>
#include "ATen/Tensor.h"
#include "utils.h"
namespace ixformer::infer {
torch::Tensor ixinfer_flash_attn_unpad_with_block_tables(
torch::Tensor& query,
torch::Tensor& key_cache,
torch::Tensor& value_cache,
torch::Tensor& out,
torch::Tensor& block_tables,
torch::Tensor& cu_seq_q,
torch::Tensor& cu_seq_k,
int64_t max_seq_q,
int64_t max_seq_k,
bool is_causal,
int64_t window_left,
int64_t window_right,
double scale,
double softcap,
bool sqrt_alibi,
const std::optional<torch::Tensor>& alibi_slopes,
const std::optional<torch::Tensor>& sinks,
std::optional<torch::Tensor>& lse);
void silu_and_mul(torch::Tensor& input, torch::Tensor& output);
torch::Tensor xllm_paged_attention(
torch::Tensor& out,
torch::Tensor& query,
torch::Tensor& key_cache,
torch::Tensor& value_cache,
int64_t num_kv_heads,
double scale,
torch::Tensor& block_tables,
torch::Tensor& context_lens,
int64_t block_size,
int64_t max_context_len,
const std::optional<torch::Tensor>& alibi_slopes,
bool causal,
int32_t window_left,
int32_t window_right,
double softcap,
bool enable_cuda_graph,
bool use_sqrt_alibi,
const std::optional<torch::Tensor>& sinks);
torch::Tensor ixformer_linear(torch::Tensor& input,
torch::Tensor& weight,
int64_t act_type,
const std::optional<torch::Tensor>& bias,
const std::optional<torch::Tensor>& out,
const std::optional<bool> persistent);
torch::Tensor ixformer_linear_ex(torch::Tensor& input,
torch::Tensor& weight,
const c10::optional<torch::Tensor>& bias,
const c10::optional<torch::Tensor>& out);
void xllm_reshape_and_cache(torch::Tensor& key,
torch::Tensor& value,
torch::Tensor& key_cache,
torch::Tensor& value_cache,
torch::Tensor& slot_mapping,
int64_t key_token_stride,
int64_t value_token_stride);
void xllm_rotary_embedding(torch::Tensor& positions,
torch::Tensor& query,
torch::Tensor& key,
int64_t head_size,
torch::Tensor& cos_sin_cache,
bool is_neox);
void residual_rms_norm(torch::Tensor& input,
torch::Tensor& residual,
torch::Tensor& weight,
torch::Tensor& output,
torch::Tensor& residual_output,
const std::optional<torch::Tensor>& fused_bias,
double alpha,
double eps,
bool is_post);
void rms_norm(torch::Tensor& input,
torch::Tensor& weight,
torch::Tensor& output,
const std::optional<torch::Tensor>& fused_bias,
double eps);
void topk_softmax(torch::Tensor& topk_weights,
torch::Tensor& topk_indices,
torch::Tensor& token_expert_indices,
torch::Tensor& gating_output,
bool renormalize);
void moe_compute_token_index_api(
torch::Tensor& topk_ids,
torch::Tensor& src_dst,
torch::Tensor& dst_src,
torch::Tensor& expert_sizes_gpu,
const c10::optional<torch::Tensor>& expert_mask,
const c10::optional<torch::Tensor>& expert_sizes_cpu,
const c10::optional<torch::Tensor>& expand_tokens_gpu,
int64_t start_expert_id,
int64_t end_expert_id,
int64_t num_experts);
void moe_expand_input(torch::Tensor outputs,
torch::Tensor inputs,
torch::Tensor dst_to_src,
const c10::optional<torch::Tensor>& src_to_dst,
int64_t dst_tokens,
int64_t expand_factor);
void moe_w16a16_group_gemm(torch::Tensor output,
torch::Tensor inputs,
torch::Tensor weights,
torch::Tensor tokens_per_experts,
const c10::optional<torch::Tensor>& dst_to_src,
const c10::optional<torch::Tensor>& bias,
std::string format,
int64_t persistent,
int64_t output_n);
void moe_output_reduce_sum(torch::Tensor outputs,
torch::Tensor inputs,
const c10::optional<torch::Tensor>& mul_weight,
const c10::optional<torch::Tensor>& mask,
const c10::optional<torch::Tensor>& extra_residual,
double scaling_factor);
} // namespace ixformer::infer

View File

@@ -0,0 +1,73 @@
/* Copyright 2025-2026 The xLLM Authors.
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.
==============================================================================*/
#include "ilu_ops_api.h"
#include "util/env_var.h"
namespace xllm::kernel::ilu {
bool gemv_conditions(const torch::Tensor& input,
const torch::Tensor& weight,
const torch::Tensor& bias,
int64_t gemv_max_batch) {
// gemv input:[m,k] weight:[n,k]
// 1. m <= gemv_max_batch
// 2. k % 32 == 0 && n % 2 == 0
// 3. bias is None
torch::Tensor input_view = input.view({-1, input.size(-1)});
torch::Tensor weight_view = weight.view({-1, weight.size(-1)});
int64_t m = input_view.size(0);
int64_t k = input_view.size(1);
int64_t n = weight_view.size(0);
if (bias.defined() == false && m <= gemv_max_batch && k % 32 == 0 &&
n % 2 == 0) {
return true;
}
return false;
}
torch::Tensor matmul(torch::Tensor a,
torch::Tensor b,
std::optional<torch::Tensor> bias) {
int64_t act_type = -1;
bool persistent = false;
std::vector<int64_t> output_shape = a.sizes().vec();
if (!output_shape.empty()) {
output_shape[output_shape.size() - 1] = b.size(0);
}
torch::Tensor output = a.new_empty(output_shape);
bool use_gemv = true;
const int64_t gemv_max_batch = 1;
const bool disable_infer_gemm_ex =
xllm::util::get_bool_env("DISABLE_INFER_GEMM_EX", false);
use_gemv =
use_gemv &&
gemv_conditions(a, b, bias.value_or(at::Tensor()), gemv_max_batch) &&
!disable_infer_gemm_ex && (act_type == -1);
if (use_gemv) {
output = infer::ixformer_linear_ex(a, b, bias, output);
} else {
output = infer::ixformer_linear(a, b, act_type, bias, output, persistent);
}
return output;
}
} // namespace xllm::kernel::ilu

View File

@@ -0,0 +1,51 @@
/* Copyright 2025-2026 The xLLM Authors.
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.
==============================================================================*/
#include "ilu_ops_api.h"
#include "utils.h"
using namespace ixformer;
namespace xllm::kernel::ilu {
void residual_layer_norm(torch::Tensor& input,
torch::Tensor& output,
std::optional<torch::Tensor>& residual,
torch::Tensor& weight,
std::optional<torch::Tensor>& bias,
std::optional<torch::Tensor>& residual_out,
double eps) {
auto residual_ = residual.value_or(torch::zeros_like(input));
torch::Tensor residual_out_ = residual_out.value_or(torch::zeros_like(input));
infer::residual_rms_norm(input,
residual_,
weight,
output,
residual_out_,
bias,
/*alpha=*/1.0,
eps,
false);
}
void rms_norm(torch::Tensor& output,
torch::Tensor& input,
torch::Tensor& weight,
double eps) {
std::optional<torch::Tensor> fused_bias = std::nullopt;
infer::rms_norm(input, weight, output, fused_bias, eps);
}
} // namespace xllm::kernel::ilu

View File

@@ -0,0 +1,31 @@
/* Copyright 2025-2026 The xLLM Authors.
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.
==============================================================================*/
#include "ilu_ops_api.h"
#include "utils.h"
namespace xllm::kernel::ilu {
void apply_rope_pos_ids_cos_sin_cache(torch::Tensor& query,
torch::Tensor& key,
torch::Tensor& cos_sin_cache,
torch::Tensor& positions,
bool interleave) {
const int64_t head_size = cos_sin_cache.size(-1);
infer::xllm_rotary_embedding(
positions, query, key, head_size, cos_sin_cache, !interleave);
}
} // namespace xllm::kernel::ilu

View File

@@ -0,0 +1,63 @@
/* Copyright 2025-2026 The xLLM Authors.
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
namespace xllm::kernel::ilu {
#undef check_tensor_contiguous
#define check_tensor_contiguous(x, type) \
TORCH_CHECK(x.scalar_type() == type); \
TORCH_CHECK(x.is_cuda()); \
TORCH_CHECK(x.is_contiguous());
#undef check_tensor_half_bf_float
#define check_tensor_half_bf_float(x) \
TORCH_CHECK(x.scalar_type() == at::ScalarType::Half || \
x.scalar_type() == at::ScalarType::Float || \
x.scalar_type() == at::ScalarType::BFloat16); \
TORCH_CHECK(x.is_cuda());
// from torchCheckMsgImpl
inline const char* ixformer_check_msg_impl(const char* msg) { return msg; }
// // If there is just 1 user-provided C-string argument, use it.
#define IXFORMER_CHECK_MSG(cond, type, ...) \
(ixformer_check_msg_impl( \
"Expected " #cond \
" to be true, but got false. " \
"(Could this error message be improved? If so, " \
"please report an enhancement request to ixformer.)", \
##__VA_ARGS__))
#define IXFORMER_CHECK(cond, ...) \
{ \
if (!(cond)) { \
std::cerr << __FILE__ << " (" << __LINE__ << ")" \
<< "-" << __FUNCTION__ << " : " \
<< IXFORMER_CHECK_MSG(cond, "", ##__VA_ARGS__) << std::endl; \
throw std::runtime_error("IXFORMER_CHECK ERROR"); \
} \
}
#undef CUINFER_CHECK
#define CUINFER_CHECK(func) \
do { \
cuinferStatus_t status = (func); \
if (status != CUINFER_STATUS_SUCCESS) { \
std::cerr << "Error in file " << __FILE__ << " on line " << __LINE__ \
<< ": " << cuinferGetErrorString(status) << std::endl; \
throw std::runtime_error("CUINFER_CHECK ERROR"); \
} \
} while (0)
} // namespace xllm::kernel::ilu