[DSA] Route the ragged prefill top-k to the v2 kernel (#35175)
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
60ff1e33a5
commit
7fd5454335
@@ -22,6 +22,7 @@
|
||||
#include <bit>
|
||||
#include <cstdint>
|
||||
#include <iterator>
|
||||
#include <limits>
|
||||
|
||||
namespace sglang {
|
||||
|
||||
@@ -64,7 +65,7 @@ struct alignas(8) PlanItem {
|
||||
};
|
||||
static_assert(sizeof(GlobalMetadata) == 2 * sizeof(int32_t) && sizeof(PlanItem) == sizeof(GlobalMetadata));
|
||||
|
||||
struct TopKLaunchParams {
|
||||
struct TopKPagedParams {
|
||||
const float* __restrict__ scores;
|
||||
const int32_t* __restrict__ seq_lens;
|
||||
const int32_t* __restrict__ page_table;
|
||||
@@ -104,12 +105,22 @@ struct TopKLaunchParams {
|
||||
}
|
||||
};
|
||||
|
||||
struct TopKRaggedParams {
|
||||
float* __restrict__ scores; // NOTE: may write
|
||||
const int32_t* __restrict__ seq_lens;
|
||||
const int32_t* __restrict__ row_starts;
|
||||
const int32_t* __restrict__ out_offsets;
|
||||
int32_t* __restrict__ topk_indices;
|
||||
int64_t score_stride;
|
||||
uint32_t topk;
|
||||
};
|
||||
|
||||
/**
|
||||
* \brief Persistent cluster kernel for the long items. It will handle long inputs.
|
||||
* The short items are handled by the separate topk_kernel.
|
||||
*/
|
||||
template <bool kPDL>
|
||||
CLUSTER_TOPK_KERNEL void topk_persistent_cluster_kernel(const __grid_constant__ TopKLaunchParams params) {
|
||||
CLUSTER_TOPK_KERNEL void topk_persistent_cluster_kernel(const __grid_constant__ TopKPagedParams params) {
|
||||
device::enable_smem_spilling();
|
||||
__shared__ impl::MaxSmem<Cluster::Smem> smem;
|
||||
const uint32_t num_cluster_items = params.global().num_cluster_items;
|
||||
@@ -159,6 +170,85 @@ SGL_DEVICE void problem_transform(TopKProblem& problem, int32_t* output_ptr) {
|
||||
for_each_item(problem.topk, [&](uint32_t tx, uint32_t i) { problem.transform_output(tx, source_index[i]); });
|
||||
}
|
||||
|
||||
/**
|
||||
* \brief Ragged (prefill) top-k: select inside a per-row window, emit indices
|
||||
* rebased onto the flattened KV.
|
||||
*
|
||||
* Row `b` selects the top-k of `scores[b][ks : ks + seq_lens[b]]` (`ks =
|
||||
* row_starts[b]`) and writes `selected_position + out_offsets[b]`, `-1` padded.
|
||||
* No page table and no plan: the DeepGEMM contiguous-KV indexer emits columns
|
||||
* that are already absolute positions in the batch's flattened KV, so an add is
|
||||
* the whole transform. One block per row -- prefill has thousands of rows, so
|
||||
* the cluster path (which exists to split ONE row across blocks) is never worth
|
||||
* it here.
|
||||
*
|
||||
* The window start is an arbitrary token offset, so the 16-byte vectorized load
|
||||
* needs the row pointer rounded down to a 4-float boundary. The <= 3 elements
|
||||
* that pulls in are columns of a preceding request -- real finite scores that
|
||||
* would otherwise win the selection -- so they are masked in place first. That
|
||||
* write races with nothing and needs no barrier of its own:
|
||||
* - one block owns the row, and a column of row `b` is read by no other row;
|
||||
* - the score buffer is dead once the top-k has run;
|
||||
* - every forward() below opens with its smem init and a `__syncthreads()`
|
||||
* before it reads any score. That barrier both publishes the mask to
|
||||
* whichever thread loads the head vector and keeps the compiler from
|
||||
* hoisting those loads above the store -- store and loads reach the same row
|
||||
* through two `__restrict__` pointers, which otherwise licenses exactly that
|
||||
* reordering.
|
||||
* It must however land after the PDL wait, or the indexer overwrites it.
|
||||
*/
|
||||
template <bool kPDL>
|
||||
TOPK_KERNEL void topk_ragged_kernel(const __grid_constant__ TopKRaggedParams params) {
|
||||
device::enable_smem_spilling();
|
||||
constexpr uint32_t kVecSize = impl::TopKStreaming::kVecSize;
|
||||
const auto bx = blockIdx.x;
|
||||
// issue all metadata prefetch ahead of time
|
||||
const auto seq_len = static_cast<uint32_t>(params.seq_lens[bx]);
|
||||
const auto offset = params.out_offsets[bx];
|
||||
const auto row_start = params.row_starts == nullptr ? 0u : params.row_starts[bx];
|
||||
const auto topk = params.topk;
|
||||
const auto out = params.topk_indices + bx * static_cast<int64_t>(topk);
|
||||
|
||||
if (seq_len <= topk) {
|
||||
device::PDLWaitPrimary<kPDL>();
|
||||
for_each_item(topk, [&](uint32_t tx, uint32_t) {
|
||||
out[tx] = tx < seq_len ? static_cast<int32_t>(tx) + offset : -1; // note: need offset
|
||||
});
|
||||
return;
|
||||
}
|
||||
|
||||
const auto rem = row_start % kVecSize;
|
||||
const auto score = params.scores + bx * params.score_stride;
|
||||
if (rem != 0) {
|
||||
// The mask has to land after the indexer has retired
|
||||
// Otherwise it may be accidentally overwritten by DG upstream
|
||||
device::PDLWaitPrimary<kPDL>();
|
||||
static_assert(kVecSize <= kBlockSize, "not enough threads ");
|
||||
if (const auto tx = threadIdx.x; tx < rem) {
|
||||
score[row_start - rem + tx] = -std::numeric_limits<float>::max();
|
||||
}
|
||||
}
|
||||
|
||||
const auto problem = TopKProblem{
|
||||
.in = score + (row_start - rem),
|
||||
.out = out,
|
||||
.page_table = nullptr, // unused
|
||||
.topk = topk,
|
||||
.seq_len = seq_len + rem,
|
||||
.page_bits = 1, // unused
|
||||
.bias = offset - static_cast<int32_t>(rem),
|
||||
};
|
||||
__shared__ impl::MaxSmem<Register2::Smem, Register4::Smem, Streaming::Smem> smem;
|
||||
if (problem.seq_len <= Register2::kMaxSeqLen) {
|
||||
Register2::forward<kPDL>(problem, &smem);
|
||||
} else if (problem.seq_len <= Register4::kMaxSeqLen) {
|
||||
Register4::forward<kPDL>(problem, &smem);
|
||||
} else {
|
||||
Streaming::forward<kPDL>(problem, &smem);
|
||||
}
|
||||
// PDL trigger secondary at the end the block typically has no use, so ignore it
|
||||
}
|
||||
|
||||
/**
|
||||
* \brief Main kernel for the short items and epilogue of long items.
|
||||
* \tparam kPDL whether to use PDL to synchronize with the cluster kernel (if any)
|
||||
@@ -169,7 +259,7 @@ SGL_DEVICE void problem_transform(TopKProblem& problem, int32_t* output_ptr) {
|
||||
* - Level 3: max_seq_len > cluster_floor -> + epilogue process of cluster path
|
||||
*/
|
||||
template <bool kPDL, int kLevel, TopKMode kMode>
|
||||
TOPK_KERNEL void topk_main_kernel(const __grid_constant__ TopKLaunchParams params) {
|
||||
TOPK_KERNEL void topk_main_kernel(const __grid_constant__ TopKPagedParams params) {
|
||||
device::enable_smem_spilling();
|
||||
auto problem = params.problem(blockIdx.x);
|
||||
constexpr uint32_t kU32Max = std::numeric_limits<uint32_t>::max();
|
||||
@@ -217,7 +307,7 @@ TOPK_KERNEL void topk_main_kernel(const __grid_constant__ TopKLaunchParams param
|
||||
}
|
||||
|
||||
template <bool kPDL, TopKMode kMode>
|
||||
CLUSTER_TOPK_KERNEL void topk_small_batch_kernel(const __grid_constant__ TopKLaunchParams params) {
|
||||
CLUSTER_TOPK_KERNEL void topk_small_batch_kernel(const __grid_constant__ TopKPagedParams params) {
|
||||
device::enable_smem_spilling();
|
||||
auto problem = params.problem(blockIdx.x);
|
||||
__shared__ impl::MaxSmem<Streaming::Smem, Cluster::Smem> smem;
|
||||
@@ -378,7 +468,7 @@ struct TopKKernel {
|
||||
static_cluster_threshold);
|
||||
}
|
||||
|
||||
static void transform(
|
||||
static void transform_paged(
|
||||
const tvm::ffi::TensorView scores,
|
||||
const tvm::ffi::TensorView seq_lens,
|
||||
const tvm::ffi::Optional<tvm::ffi::TensorView> page_table,
|
||||
@@ -444,7 +534,7 @@ struct TopKKernel {
|
||||
// The floor is chosen on the host per launch.
|
||||
constexpr uint32_t kClusterFloorSmall = 32768;
|
||||
constexpr uint32_t kSmallBatchLowFloor = 15;
|
||||
const auto params = TopKLaunchParams{
|
||||
const auto params = TopKPagedParams{
|
||||
.scores = static_cast<const float*>(scores.data_ptr()),
|
||||
.seq_lens = static_cast<const int32_t*>(seq_lens.data_ptr()),
|
||||
.page_table = page_table_ptr,
|
||||
@@ -498,6 +588,78 @@ struct TopKKernel {
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
/**
|
||||
* \brief Ragged (prefill) variant of `transform`: per-row window, additive
|
||||
* output transform, no page table and no plan.
|
||||
*
|
||||
* `scores` is written in place: the <= 3 columns the 16-byte-aligned read base
|
||||
* pulls in ahead of each row's window are masked out (see
|
||||
* `topk_ragged_kernel`). They are invalid for that row, and the buffer has no
|
||||
* consumer after this call.
|
||||
*
|
||||
* `row_starts` absent means every window starts at column 0, which is the
|
||||
* single-request case; `out_offsets` is added to every selected position and
|
||||
* is what rebases them onto the flattened KV.
|
||||
*/
|
||||
static void transform_ragged(
|
||||
const tvm::ffi::TensorView scores,
|
||||
const tvm::ffi::TensorView seq_lens,
|
||||
const tvm::ffi::Optional<tvm::ffi::TensorView> row_starts,
|
||||
const tvm::ffi::TensorView out_offsets,
|
||||
const tvm::ffi::TensorView topk_indices) {
|
||||
using namespace host;
|
||||
auto B = SymbolicSize{"batch_size"};
|
||||
auto L = SymbolicSize{"max_seq_len"};
|
||||
auto S = SymbolicSize{"score_stride"};
|
||||
auto K = SymbolicSize{"topk"};
|
||||
auto device_ = SymbolicDevice{};
|
||||
device_.set_options<kDLCUDA>();
|
||||
|
||||
TensorMatcher({B, L}) // score
|
||||
.with_strides({S, 1})
|
||||
.with_dtype<float>()
|
||||
.with_device(device_)
|
||||
.verify(scores);
|
||||
TensorMatcher({B}) // seq_lens
|
||||
.with_dtype<int32_t>()
|
||||
.with_device(device_)
|
||||
.verify(seq_lens);
|
||||
TensorMatcher({B}) // out_offsets
|
||||
.with_dtype<int32_t>()
|
||||
.with_device(device_)
|
||||
.verify(out_offsets);
|
||||
TensorMatcher({B, K}) // topk_indices
|
||||
.with_dtype<int32_t>()
|
||||
.with_device(device_)
|
||||
.verify(topk_indices);
|
||||
const int32_t* row_starts_ptr = nullptr;
|
||||
if (row_starts.has_value()) {
|
||||
TensorMatcher({B}) // row_starts
|
||||
.with_dtype<int32_t>()
|
||||
.with_device(device_)
|
||||
.verify(row_starts.value());
|
||||
row_starts_ptr = static_cast<const int32_t*>(row_starts.value().data_ptr());
|
||||
}
|
||||
|
||||
RuntimeCheck(S.unwrap() % 4 == 0, "score_stride must be a multiple of 4 (16-byte vectorized load)");
|
||||
const auto topk = static_cast<uint32_t>(K.unwrap());
|
||||
RuntimeCheck(topk > 0 && topk <= kMaxTopK, "topk must be in (0, 2048]");
|
||||
|
||||
constexpr bool kUsePDL = true;
|
||||
const auto params = TopKRaggedParams{
|
||||
.scores = static_cast<float*>(scores.data_ptr()),
|
||||
.seq_lens = static_cast<const int32_t*>(seq_lens.data_ptr()),
|
||||
.row_starts = row_starts_ptr,
|
||||
.out_offsets = static_cast<const int32_t*>(out_offsets.data_ptr()),
|
||||
.topk_indices = static_cast<int32_t*>(topk_indices.data_ptr()),
|
||||
.score_stride = S.unwrap(),
|
||||
.topk = topk,
|
||||
};
|
||||
LaunchKernel(static_cast<uint32_t>(B.unwrap()), kBlockSize, device_.unwrap())
|
||||
.config({.use_pdl = kUsePDL})
|
||||
.launch(topk_ragged_kernel<kUsePDL>, params);
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace sglang
|
||||
|
||||
@@ -176,9 +176,10 @@ struct TopKProblem {
|
||||
uint32_t topk;
|
||||
uint32_t seq_len;
|
||||
uint32_t page_bits;
|
||||
int32_t bias = 0; // needed by ragged mode
|
||||
|
||||
SGL_DEVICE void emit(uint32_t pos, uint32_t raw_idx) const {
|
||||
out[pos] = static_cast<int32_t>(raw_idx);
|
||||
out[pos] = static_cast<int32_t>(raw_idx) + bias;
|
||||
}
|
||||
SGL_DEVICE void transform_output(uint32_t t, int32_t raw) const {
|
||||
out[t] = raw < 0 ? -1 : page_to_indices(page_table, raw, page_bits);
|
||||
|
||||
@@ -39,7 +39,8 @@ def _jit_topk_v2_module():
|
||||
make_name("topk_v2"),
|
||||
cuda_files=["deepseek_v4/topk_v2.cuh"],
|
||||
cuda_wrappers=[
|
||||
("topk_transform", "TopKKernel::transform"),
|
||||
("topk_transform_paged", "TopKKernel::transform_paged"),
|
||||
("topk_transform_ragged", "TopKKernel::transform_ragged"),
|
||||
("topk_plan", "TopKKernel::plan"),
|
||||
],
|
||||
)
|
||||
@@ -86,6 +87,34 @@ def plan_topk_v2(seq_lens: torch.Tensor, static_threshold: int = 0) -> torch.Ten
|
||||
return metadata
|
||||
|
||||
|
||||
def topk_transform_ragged_v2(
|
||||
scores: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
*,
|
||||
out_offsets: torch.Tensor,
|
||||
out_indices: torch.Tensor,
|
||||
row_starts: Optional[torch.Tensor] = None,
|
||||
) -> None:
|
||||
"""Ragged (prefill) fused top-k for a contiguous-KV score matrix.
|
||||
|
||||
Row ``i`` selects the top-k of ``scores[i, ks : ks + seq_lens[i]]`` (``ks =
|
||||
row_starts[i]``, 0 when ``row_starts`` is omitted) and writes
|
||||
``selected_position + out_offsets[i]`` into ``out_indices``, ``-1`` padded.
|
||||
With the production convention ``out_offsets == row_starts`` that is the
|
||||
column index itself, i.e. the token's slot in the batch's flattened KV.
|
||||
|
||||
Unlike :func:`topk_transform_512_v2` this needs no page table and no plan
|
||||
(the cluster path only pays off for very few rows, and prefill has many).
|
||||
|
||||
IMPORTANT: ``scores`` is written in place -- the <= 3 columns ahead of each
|
||||
row's window that the 16-byte-aligned read base pulls in are masked out.
|
||||
They are invalid for that row and the buffer must have no other consumer.
|
||||
``seq_lens`` entries must be NON-NEGATIVE, as for the paged entry point.
|
||||
"""
|
||||
module = _jit_topk_v2_module()
|
||||
module.topk_transform_ragged(scores, seq_lens, row_starts, out_offsets, out_indices)
|
||||
|
||||
|
||||
def topk_transform_512_v2(
|
||||
scores: torch.Tensor,
|
||||
seq_lens: torch.Tensor,
|
||||
@@ -115,7 +144,7 @@ def topk_transform_512_v2(
|
||||
the output is all -1.
|
||||
"""
|
||||
module = _jit_topk_v2_module()
|
||||
module.topk_transform(
|
||||
module.topk_transform_paged(
|
||||
scores,
|
||||
seq_lens,
|
||||
page_tables,
|
||||
|
||||
@@ -75,6 +75,7 @@ class BaseIndexerMetadata(ABC):
|
||||
self,
|
||||
logits: torch.Tensor,
|
||||
topk: int,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Perform topk selection on the logits and possibly transform the result.
|
||||
|
||||
@@ -114,6 +114,25 @@ class DSATopKBackend(Enum):
|
||||
):
|
||||
return _topk_transform_v2_paged(logits, lengths, topk, attn_metadata)
|
||||
|
||||
# Extend-shaped RAGGED top-k for the SGL backend routes to the same v2
|
||||
# kernel through its ragged entry point: no page table (the columns are
|
||||
# already flattened-KV positions), no plan (prefill has enough rows that
|
||||
# the cluster path never applies), just a per-row window and an additive
|
||||
# output transform. `batch_idx_list` is not None only on the prefill-CP
|
||||
# path, whose `topk_indices_offset` is built from cu_seqlens_q rather
|
||||
# than the KV bases -- leave that one on the legacy kernel.
|
||||
if (
|
||||
self.should_use_topk_v2()
|
||||
and topk_transform_method == TopkTransformMethod.RAGGED
|
||||
and topk_indices_offset is not None
|
||||
and batch_idx_list is None
|
||||
and 0 < topk <= 2048
|
||||
and lengths.shape[0] == logits.shape[0] == topk_indices_offset.shape[0]
|
||||
):
|
||||
return _topk_transform_v2_ragged(
|
||||
logits, lengths, topk, topk_indices_offset, row_starts
|
||||
)
|
||||
|
||||
# The legacy transforms below read attn_metadata.page_table_1 (page_size=1),
|
||||
# which is always present here: the fold only drops it for the decode case
|
||||
# dispatched to v2 above.
|
||||
@@ -305,6 +324,37 @@ def _topk_transform_v2_paged(
|
||||
return out
|
||||
|
||||
|
||||
def _topk_transform_v2_ragged(
|
||||
logits: torch.Tensor,
|
||||
lengths: torch.Tensor,
|
||||
topk: int,
|
||||
topk_indices_offset: torch.Tensor,
|
||||
row_starts: Optional[torch.Tensor],
|
||||
) -> torch.Tensor:
|
||||
"""Fused ragged top-k via the DeepSeek-V4 v2 JIT kernel.
|
||||
|
||||
``logits`` is written in place: the kernel reads from a 16-byte-aligned base
|
||||
and masks the <= 3 columns that pulls in ahead of the window. Those columns
|
||||
belong to a preceding request of the same row, and the score buffer is dead
|
||||
after the top-k (see ``DSAIndexer._get_topk_ragged``).
|
||||
|
||||
Preconditions match the paged helper: fp32 scores with unit row stride and a
|
||||
16B-aligned row stride (DeepGEMM's contiguous-KV output satisfies this by
|
||||
construction), int32 non-negative lengths, and ``0 < topk <= 2048``.
|
||||
"""
|
||||
from sglang.kernels.ops.attention.dsv4.topk import topk_transform_ragged_v2
|
||||
|
||||
out = logits.new_empty((logits.shape[0], topk), dtype=torch.int32)
|
||||
topk_transform_ragged_v2(
|
||||
logits,
|
||||
lengths,
|
||||
out_offsets=topk_indices_offset,
|
||||
out_indices=out,
|
||||
row_starts=row_starts,
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def _build_flashinfer_paged_args(
|
||||
attn_metadata,
|
||||
row_starts: Optional[torch.Tensor],
|
||||
|
||||
@@ -5,6 +5,7 @@ from sglang.kernels.ops.attention.dsv4.topk import (
|
||||
plan_topk_v2,
|
||||
topk_transform_512,
|
||||
topk_transform_512_v2,
|
||||
topk_transform_ragged_v2,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
@@ -14,6 +15,8 @@ register_cuda_ci(
|
||||
|
||||
# Compressed page size used by the DSA indexer (real value is 256 // 4 = 64).
|
||||
PAGE_SIZE = 64
|
||||
# NOTE: currently torch baseline is disabled, since it's too slow
|
||||
DISABLE_TORCH = True
|
||||
|
||||
|
||||
def _make_inputs(batch_size: int, seq_len: int, k: int):
|
||||
@@ -44,7 +47,7 @@ def _make_p1_table(batch_size: int, seq_len: int):
|
||||
return src_page_table, lengths
|
||||
|
||||
|
||||
def _build_fn(provider: str, batch_size: int, seq_len: int, k: int):
|
||||
def _build_paged_fn(provider: str, batch_size: int, seq_len: int, k: int):
|
||||
scores, seq_lens, page_table, out = _make_inputs(batch_size, seq_len, k)
|
||||
N = PAGE_SIZE
|
||||
|
||||
@@ -72,19 +75,67 @@ def _build_fn(provider: str, batch_size: int, seq_len: int, k: int):
|
||||
return fn, (scores, seq_lens, page_table)
|
||||
|
||||
|
||||
def _build_ragged_fn(provider: str, batch_size: int, seq_len: int, k: int):
|
||||
scores, seq_lens, _, out = _make_inputs(batch_size, seq_len, k)
|
||||
offsets = torch.arange(batch_size, dtype=torch.int32, device="cuda") * seq_len
|
||||
|
||||
def fn(scores, seq_lens, offsets):
|
||||
if provider == "jit_v1":
|
||||
from sgl_kernel import fast_topk_transform_ragged_fused
|
||||
|
||||
return fast_topk_transform_ragged_fused(scores, seq_lens, offsets, k)
|
||||
elif provider == "jit_v2":
|
||||
topk_transform_ragged_v2(
|
||||
scores, seq_lens, out_offsets=offsets, out_indices=out
|
||||
)
|
||||
return out
|
||||
elif provider == "flashinfer":
|
||||
from flashinfer import top_k_ragged_transform
|
||||
|
||||
return top_k_ragged_transform(scores, offsets, seq_lens, k)
|
||||
elif provider == "torch":
|
||||
idx = scores.topk(k, dim=-1).indices.to(torch.int32) # (batch, k)
|
||||
return idx + offsets.unsqueeze(1)
|
||||
else:
|
||||
raise ValueError(f"unknown provider {provider}")
|
||||
|
||||
return fn, (scores, seq_lens, offsets)
|
||||
|
||||
|
||||
PRROVIDERS = ["jit_v1", "jit_v2", "flashinfer"]
|
||||
if not DISABLE_TORCH:
|
||||
PRROVIDERS.append("torch")
|
||||
|
||||
|
||||
@marker.parametrize("k", [512, 1024, 2048], [512])
|
||||
@marker.parametrize("seq_len", [2**x for x in range(10, 19)], [4096, 65536])
|
||||
@marker.parametrize("batch_size", [2**x for x in range(13)], [1, 128, 1024])
|
||||
@marker.benchmark("provider", ["jit_v1", "jit_v2", "flashinfer", "torch"])
|
||||
def benchmark(seq_len: int, batch_size: int, k: int, provider: str):
|
||||
@marker.benchmark("provider", PRROVIDERS)
|
||||
def benchmark_paged(seq_len: int, batch_size: int, k: int, provider: str):
|
||||
if k > seq_len:
|
||||
marker.skip("k cannot be larger than seq_len")
|
||||
if k == 2048 and provider == "jit_v1":
|
||||
marker.skip("jit_v1 does not support k=2048")
|
||||
|
||||
fn, input_args = _build_fn(provider, batch_size, seq_len, k)
|
||||
fn, input_args = _build_paged_fn(provider, batch_size, seq_len, k)
|
||||
return marker.do_bench(fn, input_args=input_args, memory_args=input_args[:2])
|
||||
|
||||
|
||||
@marker.parametrize("k", [512, 1024, 2048], [2048])
|
||||
@marker.parametrize("seq_len", [2**x for x in range(10, 19)], [4096, 65536])
|
||||
# NOTE: prefill workload should be heavier than decode; not common for short extend
|
||||
@marker.parametrize("batch_size", [2**x for x in range(7, 14)], [128, 1024])
|
||||
@marker.benchmark("provider", PRROVIDERS)
|
||||
def benchmark_ragged(seq_len: int, batch_size: int, k: int, provider: str):
|
||||
if k > seq_len:
|
||||
marker.skip("k cannot be larger than seq_len")
|
||||
if k != 2048 and provider == "jit_v1":
|
||||
marker.skip("jit_v1 (here, sgl-AOT) only support k=2048")
|
||||
|
||||
fn, input_args = _build_ragged_fn(provider, batch_size, seq_len, k)
|
||||
return marker.do_bench(fn, input_args=input_args, memory_args=input_args[:2])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
benchmark.run()
|
||||
benchmark_paged.run()
|
||||
benchmark_ragged.run()
|
||||
|
||||
@@ -29,7 +29,11 @@ import sys
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from sglang.kernels.ops.attention.dsv4.topk import plan_topk_v2, topk_transform_512_v2
|
||||
from sglang.kernels.ops.attention.dsv4.topk import (
|
||||
plan_topk_v2,
|
||||
topk_transform_512_v2,
|
||||
topk_transform_ragged_v2,
|
||||
)
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=90, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
||||
@@ -265,5 +269,112 @@ def test_topk_v2_output_indices(batch: int, seq: int, k: int) -> None:
|
||||
_assert_topk_close(scores.cpu(), ref_raw, our_raw, batch, seq_lens.cpu(), k)
|
||||
|
||||
|
||||
# --- ragged entry point ------------------------------------------------------
|
||||
# Rows select inside `[row_start, row_start + seq_len)` of their score row and
|
||||
# emit `position + offset`. The window start is an arbitrary token offset, so
|
||||
# every `row_start % 4` residue must be covered: the kernel reads from a
|
||||
# 16-byte-aligned base and masks the <=3 columns that pulls in ahead of the
|
||||
# window. Everything outside the window is filled with OUTSIDE_SCORE, which
|
||||
# beats every in-window score, so any leak shows up as a wrong selection.
|
||||
OUTSIDE_SCORE = 1e3
|
||||
|
||||
# (name, per-row (row_start, length)) spanning every template and residue.
|
||||
RAGGED_CONFIGS = [
|
||||
# one length per template band, all four residues plus aligned starts
|
||||
("trivial", [(s, 1500) for s in (0, 1, 2, 3, 4, 7, 4096, 4099)]),
|
||||
("register2", [(s, 6000) for s in (0, 1, 2, 3, 4, 7, 4096, 4099)]),
|
||||
("register4", [(s, 12000) for s in (0, 1, 2, 3, 4, 7, 4096, 4099)]),
|
||||
("streaming", [(s, 40000) for s in (0, 1, 2, 3, 4, 7, 4096, 4099)]),
|
||||
# mixed bands in one launch, laid out back to back like a real prefill batch
|
||||
(
|
||||
"mixed",
|
||||
[
|
||||
(0, 1000),
|
||||
(1000, 3000),
|
||||
(4000, 9000),
|
||||
(13000, 20000),
|
||||
(33000, 1),
|
||||
(33001, 2047),
|
||||
],
|
||||
),
|
||||
# boundaries: seq == k, seq == k + 1, and the register/streaming edges
|
||||
(
|
||||
"boundaries",
|
||||
[(1, 2048), (2049, 2049), (4098, 8192), (12290, 8193), (20483, 16385)],
|
||||
),
|
||||
("long_ctx", [(0, 131072), (131072, 65537), (196609, 100000)]),
|
||||
]
|
||||
|
||||
|
||||
def _make_ragged(rows, offset_shift, device):
|
||||
width = ((max(s + n for s, n in rows)) + 3) & ~3
|
||||
scores = torch.full(
|
||||
(len(rows), width), OUTSIDE_SCORE, dtype=torch.float32, device=device
|
||||
)
|
||||
for i, (start, length) in enumerate(rows):
|
||||
scores[i, start : start + length] = torch.randn(length, device=device)
|
||||
starts = torch.tensor([s for s, _ in rows], dtype=torch.int32, device=device)
|
||||
lengths = torch.tensor([n for _, n in rows], dtype=torch.int32, device=device)
|
||||
return scores, starts, lengths, starts + offset_shift
|
||||
|
||||
|
||||
def _run_ragged(scores, lengths, starts, offsets, k):
|
||||
"""Selected positions per row, rebased back to window-relative."""
|
||||
out = torch.empty((scores.shape[0], k), dtype=torch.int32, device=scores.device)
|
||||
topk_transform_ragged_v2(
|
||||
scores, lengths, out_offsets=offsets, out_indices=out, row_starts=starts
|
||||
)
|
||||
torch.cuda.synchronize()
|
||||
off = offsets.cpu().tolist()
|
||||
return [
|
||||
[v - off[i] for v in row if v != -1] for i, row in enumerate(out.cpu().tolist())
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("k", [512, 1024, 2048])
|
||||
@pytest.mark.parametrize("offset_shift", [0, 4321])
|
||||
@pytest.mark.parametrize("name,rows", RAGGED_CONFIGS)
|
||||
@torch.inference_mode()
|
||||
def test_topk_v2_ragged_window(name: str, rows, k: int, offset_shift: int) -> None:
|
||||
torch.manual_seed(len(rows) * 7919 + k + offset_shift)
|
||||
device = "cuda"
|
||||
scores, starts, lengths, offsets = _make_ragged(rows, offset_shift, device)
|
||||
before = scores.clone()
|
||||
|
||||
our_raw = _run_ragged(scores, lengths, starts, offsets, k)
|
||||
|
||||
# reference on the window slice, padded to a common width for the helper
|
||||
max_len = max(n for _, n in rows)
|
||||
windows = torch.zeros(len(rows), max_len, dtype=torch.float32)
|
||||
for i, (start, length) in enumerate(rows):
|
||||
windows[i, :length] = before[i, start : start + length].cpu()
|
||||
ref_raw = _reference(windows, lengths.cpu(), k)
|
||||
_assert_topk_close(windows, ref_raw, our_raw, len(rows), lengths.cpu(), k)
|
||||
|
||||
# the only legal in-place write is the <=3 masked columns ahead of a window
|
||||
# that the kernel actually reads (trivial rows read nothing)
|
||||
changed = (scores != before).cpu()
|
||||
for i, (start, length) in enumerate(rows):
|
||||
allowed = torch.zeros(scores.shape[1], dtype=torch.bool)
|
||||
if length > k:
|
||||
allowed[start - start % 4 : start] = True
|
||||
stray = (changed[i] & ~allowed).nonzero().flatten().tolist()
|
||||
assert not stray, f"row {i} ({name}) wrote outside its masked head: {stray[:8]}"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("k", [512, 2048])
|
||||
@torch.inference_mode()
|
||||
def test_topk_v2_ragged_no_row_starts(k: int) -> None:
|
||||
"""`row_starts=None` means every window starts at column 0."""
|
||||
torch.manual_seed(4242 + k)
|
||||
device = "cuda"
|
||||
rows = [(0, 900), (0, 5000), (0, 20000), (0, 70000)]
|
||||
scores, starts, lengths, offsets = _make_ragged(rows, 0, device)
|
||||
explicit = _run_ragged(scores.clone(), lengths, starts, offsets, k)
|
||||
implicit = _run_ragged(scores.clone(), lengths, None, offsets, k)
|
||||
for i in range(len(rows)):
|
||||
assert sorted(explicit[i]) == sorted(implicit[i]), f"row {i} differs"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
|
||||
Reference in New Issue
Block a user