94057c3d3e
PR Test (NPU) / check-changes (push) Has been cancelled
PR Test (NPU) / pr-gate (push) Has been cancelled
PR Test (NPU) / set-image-config (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-1-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (0) (push) Has been cancelled
PR Test (NPU) / stage-b-test-2-npu-a2 (1) (push) Has been cancelled
PR Test (NPU) / stage-b-test-4-npu-a3 (push) Has been cancelled
PR Test (NPU) / stage-b-test-16-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-1-npu-a3 (push) Has been cancelled
PR Test (NPU) / multimodal-gen-test-2-npu-a3 (push) Has been cancelled
PR Test (Arm64) / pr-gate (push) Has been cancelled
PR Test (Arm64) / check-changes (push) Has been cancelled
PR Test (Arm64) / build-test (push) Has been cancelled
PR Test (sgl-router) / gate (push) Has been cancelled
PR Test (sgl-router) / tier-1 — lint (push) Has been cancelled
PR Test (sgl-router) / tier-2 — build + test (push) Has been cancelled
PR Test (sgl-router) / tier-3 — docker (placeholder) (push) Has been cancelled
PR Test (sgl-router) / tier-3 — k8s integration (push) Has been cancelled
PR Test (sgl-router) / tier-3 — e2e (push) Has been cancelled
PR Test (sgl-router) / finish (push) Has been cancelled
PR Test (NPU) / single-node-poc (map[name:qwen3_6_27b_w8a8_1p_in64k_out1k_50ms runner:linux-aarch64-a3-2 test_case:test/registered/ascend/performance/qwen3_6_27b/test_npu_qwen3_6_27b_w8a8_1p_in64k_out1k_50ms.py test_type:perf]) (push) Has been cancelled
PR Test (NPU) / pr-test-npu-finish (push) Has been cancelled
PR Test (Xeon) / pr-gate (push) Has been cancelled
PR Test (Xeon) / check-changes (push) Has been cancelled
PR Test (Xeon) / build-test (, xeon-gnr, base-b-test-cpu) (push) Has been cancelled
PR Test (XPU) / check-changes (push) Has been cancelled
PR Test (XPU) / pr-gate (push) Has been cancelled
PR Test (XPU) / stage-a-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / wait-for-stage-a (push) Has been cancelled
PR Test (XPU) / stage-b-test-1-gpu-xpu (push) Has been cancelled
PR Test (XPU) / finish (push) Has been cancelled
CI Model Inventory / build-inventory (push) Has been cancelled
Lint / lint (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Compilation Check (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Manual Policy (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark - Request Processing (push) Has been cancelled
PR Benchmark (SMG Components) / Benchmark Summary (push) Has been cancelled
PR Test (SMG) / build-wheel (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on windows (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (x86_64 - auto) (push) Has been cancelled
PR Test (SMG) / python-unit-tests (push) Has been cancelled
PR Test (SMG) / unit-tests (push) Has been cancelled
PR Test (SMG) / benchmarks (push) Has been cancelled
PR Test (SMG) / chat-completions (push) Has been cancelled
PR Test (SMG) / chat-completions-4gpu (push) Has been cancelled
PR Test (SMG) / e2e (push) Has been cancelled
PR Test (SMG) / docker-build-test (push) Has been cancelled
PR Test (SMG) / k8s-integration (push) Has been cancelled
PR Test (SMG) / finish (push) Has been cancelled
PR Test (SMG) / summarize-benchmarks (push) Has been cancelled
Release SGLang Model Gateway Docker Image / publish (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on macos (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - auto) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (aarch64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / build on linux (x86_64 - musllinux_1_1) (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Build SDist (push) Has been cancelled
Release SGLang Model Gateway to PyPI / Upload to PyPI (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (aarch64, 12.9, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu129-matrix (x86_64, 12.9, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu129 (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (aarch64, 13.0, 3.10, arm-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / build-cu130-matrix (x86_64, 13.0, 3.10, x64-kernel-build-node) (push) Has been cancelled
Release SGLang Kernels / release-cu130 (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 700) (push) Has been cancelled
Release SGLang Kernels / build-rocm-matrix (3.10, 720) (push) Has been cancelled
Release SGLang Kernels / release-rocm700 (push) Has been cancelled
Release SGLang Kernels / release-rocm720 (push) Has been cancelled
Release SGLang Kernels / build-musa43 (43, 3.10) (push) Has been cancelled
Release SGLang Kernels / release-musa43 (push) Has been cancelled
373 lines
13 KiB
Plaintext
373 lines
13 KiB
Plaintext
/* Copyright 2025 SGLang Team. All Rights Reserved.
|
|
|
|
Licensed under the Apache License, Version 2.0 (the "License");
|
|
you may not use this file except in compliance with the License.
|
|
You may obtain a copy of the License at
|
|
|
|
http://www.apache.org/licenses/LICENSE-2.0
|
|
|
|
Unless required by applicable law or agreed to in writing, software
|
|
distributed under the License is distributed on an "AS IS" BASIS,
|
|
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
See the License for the specific language governing permissions and
|
|
limitations under the License.
|
|
==============================================================================*/
|
|
|
|
#include <ATen/core/TensorBase.h>
|
|
#include <ATen/core/TensorBody.h>
|
|
#include <c10/cuda/CUDAStream.h>
|
|
#include <c10/macros/Macros.h>
|
|
#include <c10/util/Exception.h>
|
|
#include <cuda.h>
|
|
#include <cuda_fp16.h>
|
|
|
|
#include <cstddef>
|
|
#include <cstdint>
|
|
#include <optional>
|
|
|
|
namespace {
|
|
|
|
constexpr uint32_t kMaxTopK = 1024;
|
|
constexpr uint32_t kBlockSize = 512;
|
|
|
|
#ifdef SGL_TOPK_DYNAMIC_SMEM_BYTES
|
|
constexpr size_t kSMEM = static_cast<size_t>(SGL_TOPK_DYNAMIC_SMEM_BYTES);
|
|
#else
|
|
constexpr size_t kSMEM = 48 * 1024; // bytes
|
|
#endif
|
|
static_assert(kSMEM % (2 * sizeof(int32_t)) == 0, "kSMEM must be a multiple of 8 bytes.");
|
|
|
|
struct TopKParams {
|
|
const float* __restrict__ scores;
|
|
const int32_t* __restrict__ seq_lens;
|
|
const int32_t* __restrict__ page_table;
|
|
int32_t* __restrict__ page_indices;
|
|
int32_t* __restrict__ raw_indices;
|
|
int64_t score_stride;
|
|
int64_t page_table_stride;
|
|
uint32_t page_bits;
|
|
uint32_t topk;
|
|
int64_t output_stride;
|
|
};
|
|
|
|
__device__ __forceinline__ uint8_t convert_to_uint8(float x) {
|
|
__half h = __float2half_rn(x);
|
|
uint16_t bits = __half_as_ushort(h);
|
|
uint16_t key = (bits & 0x8000) ? static_cast<uint16_t>(~bits) : static_cast<uint16_t>(bits | 0x8000);
|
|
return static_cast<uint8_t>(key >> 8);
|
|
}
|
|
|
|
__device__ __forceinline__ uint32_t convert_to_uint32(float x) {
|
|
uint32_t bits = __float_as_uint(x);
|
|
return (bits & 0x80000000u) ? ~bits : (bits | 0x80000000u);
|
|
}
|
|
|
|
__device__ __forceinline__ int32_t
|
|
page_to_slot(const int32_t* __restrict__ page_table, uint32_t i, uint32_t page_bits) {
|
|
const uint32_t mask = (1u << page_bits) - 1u;
|
|
return (page_table[i >> page_bits] << page_bits) | static_cast<int32_t>(i & mask);
|
|
}
|
|
|
|
__device__ void naive_paged_transform(
|
|
int32_t length,
|
|
uint32_t topk,
|
|
uint32_t page_bits,
|
|
const int32_t* __restrict__ page_table,
|
|
int32_t* __restrict__ page_indices_out,
|
|
int32_t* __restrict__ raw_indices_out) {
|
|
for (uint32_t i = threadIdx.x; i < topk; i += kBlockSize) {
|
|
if (i < static_cast<uint32_t>(length)) {
|
|
page_indices_out[i] = page_to_slot(page_table, i, page_bits);
|
|
if (raw_indices_out != nullptr) {
|
|
raw_indices_out[i] = static_cast<int32_t>(i);
|
|
}
|
|
} else {
|
|
page_indices_out[i] = -1;
|
|
if (raw_indices_out != nullptr) {
|
|
raw_indices_out[i] = -1;
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
__device__ void
|
|
radix_topk(const float* __restrict__ input, int32_t* __restrict__ output, uint32_t length, uint32_t topk) {
|
|
constexpr uint32_t RADIX = 256;
|
|
constexpr uint32_t BLOCK_SIZE = kBlockSize;
|
|
constexpr uint32_t SMEM_INPUT_SIZE = kSMEM / (2 * sizeof(int32_t));
|
|
|
|
alignas(128) __shared__ uint32_t _s_histogram_buf[2][RADIX + 32];
|
|
alignas(128) __shared__ uint32_t s_counter;
|
|
alignas(128) __shared__ uint32_t s_threshold_bin_id;
|
|
alignas(128) __shared__ uint32_t s_num_input[2];
|
|
alignas(128) __shared__ int32_t s_last_remain;
|
|
|
|
extern __shared__ uint32_t s_input_idx[][SMEM_INPUT_SIZE];
|
|
|
|
const uint32_t tx = threadIdx.x;
|
|
uint32_t remain_topk = topk;
|
|
auto& s_histogram = _s_histogram_buf[0];
|
|
|
|
const auto run_cumsum = [&] {
|
|
#pragma unroll 8
|
|
for (int32_t i = 0; i < 8; ++i) {
|
|
static_assert(1 << 8 == RADIX);
|
|
if (tx < RADIX) {
|
|
const auto j = 1 << i;
|
|
const auto k = i & 1;
|
|
auto value = _s_histogram_buf[k][tx];
|
|
if (tx + j < RADIX) {
|
|
value += _s_histogram_buf[k][tx + j];
|
|
}
|
|
_s_histogram_buf[k ^ 1][tx] = value;
|
|
}
|
|
__syncthreads();
|
|
}
|
|
};
|
|
|
|
// stage 1: 8bit coarse histogram
|
|
if (tx < RADIX + 1) s_histogram[tx] = 0;
|
|
__syncthreads();
|
|
for (uint32_t idx = tx; idx < length; idx += BLOCK_SIZE) {
|
|
const auto bin = convert_to_uint8(input[idx]);
|
|
::atomicAdd(&s_histogram[bin], 1);
|
|
}
|
|
__syncthreads();
|
|
run_cumsum();
|
|
if (tx < RADIX && s_histogram[tx] > remain_topk && s_histogram[tx + 1] <= remain_topk) {
|
|
s_threshold_bin_id = tx;
|
|
s_num_input[0] = 0;
|
|
s_counter = 0;
|
|
}
|
|
__syncthreads();
|
|
|
|
{
|
|
const auto threshold_bin = s_threshold_bin_id;
|
|
remain_topk -= s_histogram[threshold_bin + 1];
|
|
if (remain_topk == 0) {
|
|
for (uint32_t idx = tx; idx < length; idx += BLOCK_SIZE) {
|
|
const uint32_t bin = convert_to_uint8(input[idx]);
|
|
if (bin > threshold_bin) {
|
|
const auto pos = ::atomicAdd(&s_counter, 1);
|
|
output[pos] = static_cast<int32_t>(idx);
|
|
}
|
|
}
|
|
__syncthreads();
|
|
return;
|
|
}
|
|
__syncthreads();
|
|
if (tx < RADIX + 1) s_histogram[tx] = 0;
|
|
__syncthreads();
|
|
|
|
for (uint32_t idx = tx; idx < length; idx += BLOCK_SIZE) {
|
|
const float raw_input = input[idx];
|
|
const uint32_t bin = convert_to_uint8(raw_input);
|
|
if (bin > threshold_bin) {
|
|
const auto pos = ::atomicAdd(&s_counter, 1);
|
|
output[pos] = static_cast<int32_t>(idx);
|
|
} else if (bin == threshold_bin) {
|
|
const auto pos = ::atomicAdd(&s_num_input[0], 1);
|
|
if (C10_LIKELY(pos < SMEM_INPUT_SIZE)) {
|
|
s_input_idx[0][pos] = idx;
|
|
const auto bin32 = convert_to_uint32(raw_input);
|
|
const auto sub_bin = (bin32 >> 24) & 0xFF;
|
|
::atomicAdd(&s_histogram[sub_bin], 1);
|
|
}
|
|
}
|
|
}
|
|
__syncthreads();
|
|
}
|
|
|
|
// stage 2: refine with 8bit radix passes
|
|
#pragma unroll 4
|
|
for (int round = 0; round < 4; ++round) {
|
|
const auto r_idx = round % 2;
|
|
|
|
const auto raw_num_input = s_num_input[r_idx];
|
|
const auto num_input = raw_num_input < SMEM_INPUT_SIZE ? raw_num_input : SMEM_INPUT_SIZE;
|
|
|
|
run_cumsum();
|
|
if (tx < RADIX && s_histogram[tx] > remain_topk && s_histogram[tx + 1] <= remain_topk) {
|
|
s_threshold_bin_id = tx;
|
|
s_num_input[r_idx ^ 1] = 0;
|
|
s_last_remain = static_cast<int32_t>(remain_topk - s_histogram[tx + 1]);
|
|
}
|
|
__syncthreads();
|
|
|
|
const auto threshold_bin = s_threshold_bin_id;
|
|
remain_topk -= s_histogram[threshold_bin + 1];
|
|
|
|
if (remain_topk == 0) {
|
|
for (uint32_t i = tx; i < num_input; i += BLOCK_SIZE) {
|
|
const auto idx = s_input_idx[r_idx][i];
|
|
const auto offset = 24 - round * 8;
|
|
const auto bin = (convert_to_uint32(input[idx]) >> offset) & 0xFF;
|
|
if (bin > threshold_bin) {
|
|
const auto pos = ::atomicAdd(&s_counter, 1);
|
|
output[pos] = static_cast<int32_t>(idx);
|
|
}
|
|
}
|
|
__syncthreads();
|
|
break;
|
|
}
|
|
__syncthreads();
|
|
if (tx < RADIX + 1) s_histogram[tx] = 0;
|
|
__syncthreads();
|
|
for (uint32_t i = tx; i < num_input; i += BLOCK_SIZE) {
|
|
const auto idx = s_input_idx[r_idx][i];
|
|
const auto raw_input = input[idx];
|
|
const auto offset = 24 - round * 8;
|
|
const auto bin = (convert_to_uint32(raw_input) >> offset) & 0xFF;
|
|
if (bin > threshold_bin) {
|
|
const auto pos = ::atomicAdd(&s_counter, 1);
|
|
output[pos] = static_cast<int32_t>(idx);
|
|
} else if (bin == threshold_bin) {
|
|
if (round == 3) {
|
|
const auto pos = ::atomicAdd(&s_last_remain, -1);
|
|
if (pos > 0) {
|
|
output[topk - pos] = static_cast<int32_t>(idx);
|
|
}
|
|
} else {
|
|
const auto pos = ::atomicAdd(&s_num_input[r_idx ^ 1], 1);
|
|
if (C10_LIKELY(pos < SMEM_INPUT_SIZE)) {
|
|
s_input_idx[r_idx ^ 1][pos] = idx;
|
|
const auto bin32 = convert_to_uint32(raw_input);
|
|
const auto sub_bin = (bin32 >> (offset - 8)) & 0xFF;
|
|
::atomicAdd(&s_histogram[sub_bin], 1);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
__syncthreads();
|
|
}
|
|
}
|
|
|
|
__global__ __launch_bounds__(kBlockSize) void deepseek_v4_topk_transform_kernel(const TopKParams params) {
|
|
const auto bid = blockIdx.x;
|
|
const auto seq_len = params.seq_lens[bid];
|
|
const auto topk = params.topk;
|
|
const auto score_ptr = params.scores + bid * params.score_stride;
|
|
const auto page_ptr = params.page_table + bid * params.page_table_stride;
|
|
const auto indices_ptr = params.page_indices + bid * params.output_stride;
|
|
const auto raw_indices_ptr =
|
|
params.raw_indices != nullptr ? params.raw_indices + bid * params.output_stride : nullptr;
|
|
|
|
if (seq_len <= static_cast<int32_t>(topk)) {
|
|
naive_paged_transform(seq_len, topk, params.page_bits, page_ptr, indices_ptr, raw_indices_ptr);
|
|
return;
|
|
}
|
|
|
|
__shared__ int32_t s_topk_indices[kMaxTopK];
|
|
radix_topk(score_ptr, s_topk_indices, static_cast<uint32_t>(seq_len), topk);
|
|
|
|
__syncthreads();
|
|
for (uint32_t i = threadIdx.x; i < topk; i += kBlockSize) {
|
|
const auto raw = s_topk_indices[i];
|
|
indices_ptr[i] = page_to_slot(page_ptr, static_cast<uint32_t>(raw), params.page_bits);
|
|
if (raw_indices_ptr != nullptr) {
|
|
raw_indices_ptr[i] = raw;
|
|
}
|
|
}
|
|
}
|
|
|
|
template <auto* f, size_t kMaxDynamicSMEM>
|
|
void setup_kernel_smem_once() {
|
|
[[maybe_unused]]
|
|
static const auto result = [] {
|
|
#ifdef USE_ROCM
|
|
return ::cudaFuncSetAttribute(
|
|
reinterpret_cast<const void*>(f), ::cudaFuncAttributeMaxDynamicSharedMemorySize, kMaxDynamicSMEM);
|
|
#else
|
|
return ::cudaFuncSetAttribute(f, ::cudaFuncAttributeMaxDynamicSharedMemorySize, kMaxDynamicSMEM);
|
|
#endif
|
|
}();
|
|
TORCH_CHECK(
|
|
result == cudaSuccess, "deepseek_v4_topk_transform: cudaFuncSetAttribute failed: ", ::cudaGetErrorString(result));
|
|
}
|
|
|
|
} // namespace
|
|
|
|
#define CHECK_CUDA(x) TORCH_CHECK(x.is_cuda(), #x " must be a CUDA tensor")
|
|
|
|
void deepseek_v4_topk_transform_512(
|
|
const at::Tensor& scores,
|
|
const at::Tensor& seq_lens,
|
|
const at::Tensor& page_table,
|
|
at::Tensor& page_indices,
|
|
int64_t page_size,
|
|
std::optional<at::Tensor> raw_indices_opt) {
|
|
CHECK_CUDA(scores);
|
|
CHECK_CUDA(seq_lens);
|
|
CHECK_CUDA(page_table);
|
|
CHECK_CUDA(page_indices);
|
|
if (raw_indices_opt.has_value()) {
|
|
CHECK_CUDA(raw_indices_opt.value());
|
|
}
|
|
|
|
TORCH_CHECK(
|
|
scores.dim() == 2 && scores.scalar_type() == at::kFloat, "scores must be float32 with shape [B, max_seq_len]");
|
|
TORCH_CHECK(scores.stride(1) == 1, "scores must be contiguous along the last dim");
|
|
|
|
TORCH_CHECK(
|
|
seq_lens.dim() == 1 && seq_lens.is_contiguous() && seq_lens.scalar_type() == at::kInt,
|
|
"seq_lens must be int32 with shape [B], contiguous");
|
|
|
|
TORCH_CHECK(
|
|
page_table.dim() == 2 && page_table.scalar_type() == at::kInt,
|
|
"page_table must be int32 with shape [B, num_pages]");
|
|
TORCH_CHECK(page_table.stride(1) == 1, "page_table must be contiguous along the last dim");
|
|
|
|
const auto topk = page_indices.size(1);
|
|
TORCH_CHECK(
|
|
page_indices.dim() == 2 && page_indices.is_contiguous() && page_indices.scalar_type() == at::kInt,
|
|
"page_indices must be int32 with shape [B, topk], contiguous");
|
|
TORCH_CHECK(
|
|
topk > 0 && topk <= static_cast<int64_t>(kMaxTopK),
|
|
"page_indices last dim must be in [1, ",
|
|
kMaxTopK,
|
|
"], got ",
|
|
topk);
|
|
|
|
const auto B = scores.size(0);
|
|
TORCH_CHECK(
|
|
seq_lens.size(0) == B && page_table.size(0) == B && page_indices.size(0) == B,
|
|
"batch sizes must match across scores, seq_lens, page_table, page_indices");
|
|
|
|
TORCH_CHECK(
|
|
page_size > 0 && (page_size & (page_size - 1)) == 0, "page_size must be a positive power of 2, got ", page_size);
|
|
const auto page_bits = static_cast<uint32_t>(__builtin_ctzll(static_cast<unsigned long long>(page_size)));
|
|
|
|
int32_t* raw_ptr = nullptr;
|
|
if (raw_indices_opt.has_value()) {
|
|
auto& raw = raw_indices_opt.value();
|
|
TORCH_CHECK(
|
|
raw.dim() == 2 && raw.is_contiguous() && raw.scalar_type() == at::kInt,
|
|
"raw_indices must be int32 with shape [B, topk], contiguous");
|
|
TORCH_CHECK(raw.size(0) == B && raw.size(1) == topk, "raw_indices shape must match page_indices [B, ", topk, "]");
|
|
raw_ptr = raw.data_ptr<int32_t>();
|
|
}
|
|
|
|
const TopKParams params{
|
|
.scores = scores.data_ptr<float>(),
|
|
.seq_lens = seq_lens.data_ptr<int32_t>(),
|
|
.page_table = page_table.data_ptr<int32_t>(),
|
|
.page_indices = page_indices.data_ptr<int32_t>(),
|
|
.raw_indices = raw_ptr,
|
|
.score_stride = scores.stride(0),
|
|
.page_table_stride = page_table.stride(0),
|
|
.page_bits = page_bits,
|
|
.topk = static_cast<uint32_t>(topk),
|
|
.output_stride = topk,
|
|
};
|
|
|
|
const auto stream = at::cuda::getCurrentCUDAStream().stream();
|
|
const dim3 grid(static_cast<uint32_t>(B));
|
|
const dim3 block(kBlockSize);
|
|
|
|
setup_kernel_smem_once<deepseek_v4_topk_transform_kernel, kSMEM>();
|
|
deepseek_v4_topk_transform_kernel<<<grid, block, kSMEM, stream>>>(params);
|
|
|
|
const auto err = cudaGetLastError();
|
|
TORCH_CHECK(err == cudaSuccess, "deepseek_v4_topk_transform kernel launch failed: ", ::cudaGetErrorString(err));
|
|
}
|