Files
wehub-resource-sync 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
chore: import upstream snapshot with attribution
2026-07-13 12:38:16 +08:00

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));
}