Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
186 changes: 109 additions & 77 deletions python/sglang/kernels/jit/csrc/deepseek_v4/topk_v2.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,11 @@ namespace sglang {
namespace impl = device::topk;
using impl::TopKProblem;

enum class TopKMode {
INDICES, ///< raw selected indices into `out`; `page_table` unused
PAGE_TABLE, ///< page-table-transformed indices into `out`
};

using Register2 = impl::TopKRegister<2>; // <= 8192, register-resident, 1 read
using Register4 = impl::TopKRegister<4>; // <= 16384, register-resident, 1 read
using Streaming = impl::TopKStreaming;
Expand Down Expand Up @@ -64,7 +69,6 @@ struct TopKLaunchParams {
const int32_t* __restrict__ seq_lens;
const int32_t* __restrict__ page_table;
int32_t* __restrict__ page_indices;
int32_t* __restrict__ raw_indices; // optional raw (pre-transform) indices output; nullptr if unused
const PlanItem* __restrict__ metadata; // [0]=GlobalMetadata, [1+i]=PlanItem
int64_t score_stride;
int64_t page_table_stride;
Expand All @@ -89,7 +93,6 @@ struct TopKLaunchParams {
return TopKProblem{
.in = scores + batch_id * score_stride,
.out = page_indices + batch_id * k,
.raw_out = raw_indices != nullptr ? raw_indices + batch_id * k : nullptr,
.page_table = page_table + batch_id * page_table_stride,
.topk = topk,
.seq_len = seq_len,
Expand Down Expand Up @@ -133,12 +136,17 @@ SGL_DEVICE void for_each_item(uint32_t topk, const F& f) {
}
}

template <bool kPDL>
template <bool kPDL, TopKMode kMode>
SGL_DEVICE void trivial_transform(const TopKProblem& problem) {
device::PDLWaitPrimary<kPDL>();
device::PDLTriggerSecondary<kPDL>();
for_each_item(problem.topk, [&](uint32_t tx, uint32_t) {
problem.transform_output(tx, tx < problem.seq_len ? static_cast<int32_t>(tx) : -1);
const auto idx = tx < problem.seq_len ? static_cast<int32_t>(tx) : -1;
if constexpr (kMode == TopKMode::INDICES) {
problem.emit(tx, idx);
} else {
problem.transform_output(tx, idx);
}
});
}

Expand All @@ -160,17 +168,25 @@ SGL_DEVICE void problem_transform(TopKProblem& problem, int32_t* output_ptr) {
* - Level 2: max_seq_len <= cluster_floor -> trivial + register<4> + streaming
* - Level 3: max_seq_len > cluster_floor -> + epilogue process of cluster path
*/
template <bool kPDL, int kLevel>
template <bool kPDL, int kLevel, TopKMode kMode>
TOPK_KERNEL void topk_main_kernel(const __grid_constant__ TopKLaunchParams params) {
device::enable_smem_spilling();
auto problem = params.problem(blockIdx.x);
constexpr uint32_t kU32Max = std::numeric_limits<uint32_t>::max();
constexpr bool kHandleCluster = (kLevel == 3);
// Only the cluster path consumes the cluster kernel's output, so only it waits
// on that kernel (kPDLFinal). Every other path waits at most on the indexer
// (kPDLEarly) and must not be held on an SM slot until the long-running
// persistent pool retires -- that would serialize the short items behind it.
constexpr bool kPDLEarly = kPDL && !kHandleCluster;
constexpr bool kPDLFinal = kPDL && kHandleCluster;
__shared__ impl::MaxSmem<Register2::Smem, Register4::Smem, Streaming::Smem> smem;
if (problem.seq_len <= problem.topk) return trivial_transform<kPDL>(problem);
__shared__ int32_t topk_indices[kMaxTopK];
problem.out = topk_indices;
if (problem.seq_len <= problem.topk) return trivial_transform<kPDLEarly, kMode>(problem);

constexpr bool kNeedStaging = kMode != TopKMode::INDICES;
__shared__ int32_t s_topk_indices[kNeedStaging ? kMaxTopK : 1];
if constexpr (kNeedStaging) problem.out = s_topk_indices;

constexpr bool kHandleCluster = (kLevel == 3);
// non-trivial path: dispatch based on level and seq_len
const auto cluster_threshold = kHandleCluster ? params.cluster_threshold() : kU32Max;
if constexpr (kLevel == 0) {
Expand All @@ -181,34 +197,35 @@ TOPK_KERNEL void topk_main_kernel(const __grid_constant__ TopKLaunchParams param
Register4::forward<kPDL>(problem, &smem); // max_seq_len <= 16384 guarantees seq <= 16384
} else {
static_assert(kLevel == 2 || kLevel == 3, "we only support level = 0,1,2,3 now");
// if using cluster, we can delay the PDL wait
constexpr bool kPDLEarly = kPDL && !kHandleCluster;
constexpr bool kPDLFinal = kPDL && kHandleCluster;
if (problem.seq_len <= kReg4MaxSeqLen) {
Register4::forward<kPDLEarly>(problem, &smem);
} else if (problem.seq_len <= cluster_threshold) {
Streaming::forward<kPDLEarly>(problem, &smem);
} else { // cluster path do nothing here
} else {
// Cluster path: the pool already selected into our output row; the only
// work left is the epilogue, so this is the one path that waits for it.
problem.out = params.get_output_ptr(blockIdx.x);
device::PDLWaitPrimary<kPDLFinal>();
}
device::PDLWaitPrimary<kPDLFinal>();
}

// page-table transform pass (gathers kept out of the hot scatter loop),
// then trigger the dependent kernel only after the full output is written.
device::PDLTriggerSecondary<kPDL>();
__syncthreads();
problem_transform(problem, params.get_output_ptr(blockIdx.x));
if constexpr (kNeedStaging) {
__syncthreads();
problem_transform(problem, params.get_output_ptr(blockIdx.x));
}
}

template <bool kPDL>
template <bool kPDL, TopKMode kMode>
CLUSTER_TOPK_KERNEL void topk_small_batch_kernel(const __grid_constant__ TopKLaunchParams params) {
device::enable_smem_spilling();
auto problem = params.problem(blockIdx.x);
__shared__ impl::MaxSmem<Streaming::Smem, Cluster::Smem> smem;
if (problem.seq_len <= problem.topk) return trivial_transform<kPDL>(problem);
__shared__ int32_t topk_indices[kMaxTopK];
problem.out = topk_indices;
if (problem.seq_len <= problem.topk) return trivial_transform<kPDL, kMode>(problem);

constexpr bool kNeedStaging = kMode != TopKMode::INDICES;
__shared__ int32_t s_topk_indices[kNeedStaging ? kMaxTopK : 1];
if constexpr (kNeedStaging) problem.out = s_topk_indices;

// randomly elect one worker rank to avoid workload imbalance
const auto worker_rank = blockIdx.x % kClusterSize;
Expand All @@ -217,30 +234,34 @@ CLUSTER_TOPK_KERNEL void topk_small_batch_kernel(const __grid_constant__ TopKLau
if (problem.seq_len <= kReg4MaxSeqLen) {
if (blockIdx.y != worker_rank) return;
Register4::forward<kPDL>(problem, &smem);
device::PDLWaitPrimary<kPDL>();
__syncthreads();
} else if (problem.seq_len <= params.cluster_floor) {
if (blockIdx.y != worker_rank) return;
Streaming::forward<kPDL>(problem, &smem);
device::PDLWaitPrimary<kPDL>();
__syncthreads();
} else {
auto cluster = cooperative_groups::this_cluster();
problem.out = cluster.map_shared_rank(topk_indices, worker_rank);
Cluster::forward<kPDL>(problem, &smem); // write to peer's output shared memory
device::PDLWaitPrimary<kPDL>();
cluster.sync();
if (blockIdx.y != worker_rank) return;
if constexpr (kNeedStaging) {
problem.out = cluster.map_shared_rank(s_topk_indices, worker_rank);
}
Cluster::forward<kPDL>(problem, &smem);
if constexpr (kNeedStaging) {
cluster.sync();
if (blockIdx.y != worker_rank) return;
}
}

// Only the elected worker reaches here, and it mapped `topk_indices` to
// itself, so `problem.out` is this block's own buffer. Stating that keeps the
// shared::cluster address out of the load problem_transform issues -- which is
// load-bearing, not an optimization: without it cicc segfaults on CUDA 13.1+
// for sm_90a (issue #32830, previously worked around by copying `problem` in
// #32910). Verified: dropping this line reproduces the crash on 13.1/13.2/13.3.
__builtin_assume(problem.out == topk_indices);
problem_transform(problem, params.get_output_ptr(blockIdx.x));
device::PDLTriggerSecondary<kPDL>();
if constexpr (kNeedStaging) {
// Only the elected worker reaches here, and it mapped `topk_indices` to
// itself, so `problem.out` is this block's own buffer. Stating that keeps the
// shared::cluster address out of the load problem_transform issues -- which is
// load-bearing, not an optimization: without it cicc segfaults on CUDA 13.1+
// for sm_90a (issue #32830, previously worked around by copying `problem` in
// #32910). Verified: dropping this line reproduces the crash on 13.1/13.2/13.3.
__builtin_assume(problem.out == s_topk_indices);
problem_transform(problem, params.get_output_ptr(blockIdx.x));
}
}

// --- Plan: choose cluster_threshold from the seq_len distribution -----------
Expand Down Expand Up @@ -360,11 +381,10 @@ struct TopKKernel {
static void transform(
const tvm::ffi::TensorView scores,
const tvm::ffi::TensorView seq_lens,
const tvm::ffi::TensorView page_table,
const tvm::ffi::Optional<tvm::ffi::TensorView> page_table,
const tvm::ffi::TensorView page_indices,
const uint32_t page_size,
const tvm::ffi::TensorView metadata,
const tvm::ffi::Optional<tvm::ffi::TensorView> raw_indices) {
const tvm::ffi::TensorView metadata) {
using namespace host;
auto B = SymbolicSize{"batch_size"};
auto Bp1 = SymbolicSize{"batch_size_plus_1"};
Expand All @@ -384,11 +404,19 @@ struct TopKKernel {
.with_dtype<int32_t>()
.with_device(device_)
.verify(seq_lens);
TensorMatcher({B, -1}) // page_table
.with_strides({P, 1})
.with_dtype<int32_t>()
.with_device(device_)
.verify(page_table);
// Absent means "no page transform": `page_indices` then receives the raw
// selected indices and nothing dereferences a page table.
const int32_t* page_table_ptr = nullptr;
int64_t page_table_stride = 0;
if (page_table.has_value()) {
TensorMatcher({B, -1}) // page_table
.with_strides({P, 1})
.with_dtype<int32_t>()
.with_device(device_)
.verify(page_table.value());
page_table_ptr = static_cast<const int32_t*>(page_table.value().data_ptr());
page_table_stride = P.unwrap();
}
TensorMatcher({B, K}) // page_indices
.with_dtype<int32_t>()
.with_device(device_)
Expand All @@ -398,12 +426,6 @@ struct TopKKernel {
.with_device(device_)
.verify(metadata);

int32_t* raw_indices_ptr = nullptr;
if (raw_indices.has_value()) {
TensorMatcher({B, K}).with_dtype<int32_t>().with_device(device_).verify(raw_indices.value());
raw_indices_ptr = static_cast<int32_t*>(raw_indices.value().data_ptr());
}

RuntimeCheck(std::has_single_bit(page_size), "page_size must be power of 2");
RuntimeCheck(S.unwrap() % 4 == 0, "score_stride must be a multiple of 4 (16-byte vectorized load)");
RuntimeCheck(Bp1.unwrap() == B.unwrap() + 1, "invalid metadata shape");
Expand All @@ -425,46 +447,56 @@ struct TopKKernel {
const auto params = TopKLaunchParams{
.scores = static_cast<const float*>(scores.data_ptr()),
.seq_lens = static_cast<const int32_t*>(seq_lens.data_ptr()),
.page_table = static_cast<const int32_t*>(page_table.data_ptr()),
.page_table = page_table_ptr,
.page_indices = static_cast<int32_t*>(page_indices.data_ptr()),
.raw_indices = raw_indices_ptr,
.metadata = static_cast<const PlanItem*>(metadata.data_ptr()),
.score_stride = S.unwrap(),
.page_table_stride = P.unwrap(),
.page_table_stride = page_table_stride,
.topk = topk,
.page_bits = page_bits,
.cluster_floor = (batch_size <= kSmallBatchLowFloor) ? kClusterFloorSmall : kClusterFloor,
};

const bool use_cluster = (max_seq_len > params.cluster_floor) && (batch_size <= kClusterMaxBatch);
constexpr bool kUsePDL = true;
if (use_cluster) {
if (batch_size <= kNumPersistentClusters) {
LaunchKernel({batch_size, kClusterSize}, kBlockSize, device)
.config({.use_pdl = kUsePDL, .cluster_dim = dim3{1, kClusterSize}})
.launch(topk_small_batch_kernel<kUsePDL>, params);
const auto mode = page_table.has_value() ? TopKMode::PAGE_TABLE : TopKMode::INDICES;
const auto dispatch = [&]<typename F>(F&& f) {
switch (mode) {
case TopKMode::INDICES:
return f.template operator()<TopKMode::INDICES>();
default:
return f.template operator()<TopKMode::PAGE_TABLE>();
}
};
dispatch([&]<TopKMode kMode>() {
if (use_cluster) {
if (batch_size <= kNumPersistentClusters) {
LaunchKernel({batch_size, kClusterSize}, kBlockSize, device)
.config({.use_pdl = kUsePDL, .cluster_dim = dim3{1, kClusterSize}})
.launch(topk_small_batch_kernel<kUsePDL, kMode>, params);
} else {
const uint32_t num_clusters = std::min(batch_size, kNumPersistentClusters);
LaunchKernel({num_clusters, kClusterSize}, kBlockSize, device)
.config({.use_pdl = kUsePDL, .cluster_dim = dim3{1, kClusterSize}})
.launch(topk_persistent_cluster_kernel<kUsePDL>, params);
LaunchKernel(batch_size, kBlockSize, device)
.config({.use_pdl = kUsePDL})
.launch(topk_main_kernel<kUsePDL, /*kLevel=*/3, kMode>, params);
}
} else if (max_seq_len <= kReg2MaxSeqLen) {
LaunchKernel(batch_size, kBlockSize, device)
.config({.use_pdl = kUsePDL})
.launch(topk_main_kernel<kUsePDL, /*kLevel=*/0, kMode>, params);
} else if (max_seq_len <= kReg4MaxSeqLen) {
LaunchKernel(batch_size, kBlockSize, device)
.config({.use_pdl = kUsePDL})
.launch(topk_main_kernel<kUsePDL, /*kLevel=*/1, kMode>, params);
} else {
const uint32_t num_clusters = std::min(batch_size, kNumPersistentClusters);
LaunchKernel({num_clusters, kClusterSize}, kBlockSize, device)
.config({.use_pdl = kUsePDL, .cluster_dim = dim3{1, kClusterSize}})
.launch(topk_persistent_cluster_kernel<kUsePDL>, params);
LaunchKernel(batch_size, kBlockSize, device)
.config({.use_pdl = kUsePDL})
.launch(topk_main_kernel<kUsePDL, /*kLevel=*/3>, params);
.launch(topk_main_kernel<kUsePDL, /*kLevel=*/2, kMode>, params);
}
} else if (max_seq_len <= kReg2MaxSeqLen) {
LaunchKernel(batch_size, kBlockSize, device)
.config({.use_pdl = kUsePDL})
.launch(topk_main_kernel<kUsePDL, /*kLevel=*/0>, params);
} else if (max_seq_len <= kReg4MaxSeqLen) {
LaunchKernel(batch_size, kBlockSize, device)
.config({.use_pdl = kUsePDL})
.launch(topk_main_kernel<kUsePDL, /*kLevel=*/1>, params);
} else {
LaunchKernel(batch_size, kBlockSize, device)
.config({.use_pdl = kUsePDL})
.launch(topk_main_kernel<kUsePDL, /*kLevel=*/2>, params);
}
});
}
};

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -168,25 +168,19 @@ SGL_DEVICE int32_t page_to_indices(const int32_t* __restrict__ page_table, uint3

/// One batch element's worth of work. `emit(pos, raw_idx)` writes the selected raw
/// index to output slot `pos`; `transform_output` then applies the page-table
/// transform in a separate pass (and records the raw index in `raw_out` if set).
/// transform in a separate pass.
struct TopKProblem {
const float* __restrict__ in;
int32_t* __restrict__ out; // page_indices [topk]
int32_t* __restrict__ raw_out; // optional raw (pre-transform) indices [topk]; nullptr if unused
int32_t* __restrict__ out; // page_indices [topk]
const int32_t* __restrict__ page_table;
uint32_t topk;
uint32_t seq_len;
uint32_t page_bits;

// Write the raw selected index; the page-table transform is applied afterwards
// by transform_output() in a separate, pipelined pass. Keeping the per-element
// page_table gather off the atomic-serialized scatter loop is measurably faster
// for both short and long context.
SGL_DEVICE void emit(uint32_t pos, uint32_t raw_idx) const {
out[pos] = static_cast<int32_t>(raw_idx);
}
SGL_DEVICE void transform_output(uint32_t t, int32_t raw) const {
if (raw_out != nullptr) raw_out[t] = raw;
out[t] = raw < 0 ? -1 : page_to_indices(page_table, raw, page_bits);
}
};
Expand Down
15 changes: 11 additions & 4 deletions python/sglang/kernels/ops/attention/dsv4/topk.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,13 +89,21 @@ def plan_topk_v2(seq_lens: torch.Tensor, static_threshold: int = 0) -> torch.Ten
def topk_transform_512_v2(
scores: torch.Tensor,
seq_lens: torch.Tensor,
page_tables: torch.Tensor,
page_tables: Optional[torch.Tensor],
out_page_indices: torch.Tensor,
page_size: int,
metadata: torch.Tensor,
out_raw_indices: Optional[torch.Tensor] = None,
) -> None:
"""Fused top-k + page-table transform (DeepSeek-V4 top-k v2 kernel).
"""Fused top-k + optional page-table transform (DeepSeek-V4 top-k v2 kernel).

Two output modes, chosen by whether ``page_tables`` is given and resolved to
a device-side template parameter, so an unused page-table gather is compiled
out rather than skipped at runtime:

* ``page_tables=None`` -- ``out_page_indices`` receives the raw selected
indices and no page table is read.
* ``page_tables`` given -- ``out_page_indices`` receives the page-table
transform of them.

IMPORTANT: every entry of ``seq_lens`` must be NON-NEGATIVE, and
``metadata`` must come from :func:`plan_topk_v2` over the same ``seq_lens``
Expand All @@ -114,5 +122,4 @@ def topk_transform_512_v2(
out_page_indices,
page_size,
metadata,
out_raw_indices,
)
6 changes: 2 additions & 4 deletions python/sglang/srt/layers/attention/dsa/dsa_topk_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -289,8 +289,6 @@ def _topk_transform_v2_paged(
assert 0 < topk <= 2048, f"v2 top-k supports 0 < topk <= 2048, got {topk=}"

page_table = attn_metadata.real_page_table
assert page_table.dtype == torch.int32
lengths_i32 = lengths.to(torch.int32)

# The plan is preprocessed once per forward (DSAMetadata.topk_v2_plan,
# refreshed in-place under CUDA graph) and reused across layers. A missing or
Expand All @@ -302,8 +300,8 @@ def _topk_transform_v2_paged(
), "topk_v2_plan must be preprocessed per forward (see DSAMetadata.topk_v2_plan)"

page_size = attn_metadata.page_size
out = logits.new_full((num_rows, topk), -1, dtype=torch.int32)
topk_transform_512_v2(logits, lengths_i32, page_table, out, page_size, plan)
out = logits.new_empty((num_rows, topk), dtype=torch.int32)
topk_transform_512_v2(logits, lengths, page_table, out, page_size, plan)
return out


Expand Down
Loading
Loading