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

137 lines
4.9 KiB
Python

"""Test for the kpool top-k transform JIT kernel.
Ported from the former AOT sgl-kernel test (sgl-kernel/tests/test_topk.py).
The kernel selects pool groups at pool granularity, expands each selected group
to ``pool_size`` token indices, and optionally transforms those token indices
through a page table or a ragged offset.
"""
from __future__ import annotations
import sys
from typing import Optional
import pytest
import torch
from sglang.jit_kernel.kpool_topk_transform import fast_kpool_topk_transform_fused
from sglang.test.ci.ci_register import register_cuda_ci
register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="1-gpu-large")
def _ref_torch_kpool_transform_impl(
score: torch.Tensor,
lengths: torch.Tensor,
pool_size: int,
topk: int,
page_table: Optional[torch.Tensor] = None,
topk_indices_offset: Optional[torch.Tensor] = None,
seq_lens: Optional[torch.Tensor] = None,
) -> torch.Tensor:
rows = score.shape[0]
group_topk = topk // pool_size
offsets = torch.arange(pool_size, dtype=torch.int32, device=score.device)
out_cols = topk + (pool_size - 1 if seq_lens is not None else 0)
out = torch.full((rows, out_cols), -1, dtype=torch.int32, device=score.device)
for i in range(rows):
length = int(lengths[i].item())
valid_count = min(length, group_topk)
write_pos = 0
if valid_count == 0:
token_ids = torch.empty((0,), dtype=torch.int32, device=score.device)
elif length <= group_topk:
selected = torch.arange(length, dtype=torch.int32, device=score.device)
token_ids = (selected.unsqueeze(1) * pool_size + offsets).reshape(-1)
else:
selected = torch.topk(
score[i, :length], group_topk, dim=-1, sorted=False
).indices.to(torch.int32)
token_ids = (selected.unsqueeze(1) * pool_size + offsets).reshape(-1)
if token_ids.numel() > 0:
if page_table is not None:
token_ids = page_table[i, token_ids.long()].to(torch.int32)
elif topk_indices_offset is not None:
token_ids = token_ids + topk_indices_offset[i].to(torch.int32)
write_pos = valid_count * pool_size
out[i, :write_pos] = token_ids[:write_pos]
if seq_lens is not None:
tail_count = int(seq_lens[i].item()) % pool_size
if tail_count > 0:
raw_tail = length * pool_size + torch.arange(
tail_count, dtype=torch.int32, device=score.device
)
if page_table is not None:
tail = page_table[i, raw_tail.long()].to(torch.int32)
elif topk_indices_offset is not None:
tail = raw_tail + topk_indices_offset[i].to(torch.int32)
else:
tail = raw_tail
out[i, write_pos : write_pos + tail_count] = tail
return out
@pytest.mark.parametrize(
"pool_size,group_topk",
[(16, 128), (16, 160), (16, 192), (16, 224), (8, 256), (4, 512)],
)
@pytest.mark.parametrize("mode", ["raw", "paged", "ragged"])
@pytest.mark.parametrize("append_tail", [False, True])
@torch.inference_mode()
def test_kpool_topk_transform_kernel(
pool_size: int, group_topk: int, mode: str, append_tail: bool
) -> None:
torch.manual_seed(42)
bs = 17
topk = pool_size * group_topk
num_groups = 4096
score = torch.randn(bs, num_groups, dtype=torch.float32, device="cuda")
lengths = torch.randint(
group_topk + 1, num_groups + 1, (bs,), dtype=torch.int32, device="cuda"
)
page_table = None
topk_indices_offset = None
seq_lens = None
tail_counts = torch.randint(0, pool_size, (bs,), dtype=torch.int32, device="cuda")
if append_tail:
seq_lens = lengths * pool_size + tail_counts
if mode == "paged":
page_table = torch.arange(
bs * (num_groups * pool_size + pool_size),
dtype=torch.int32,
device="cuda",
).view(bs, num_groups * pool_size + pool_size)
elif mode == "ragged":
topk_indices_offset = torch.randint(
0, 2048, (bs,), dtype=torch.int32, device="cuda"
)
out_ref = _ref_torch_kpool_transform_impl(
score,
lengths,
pool_size,
topk,
page_table=page_table,
topk_indices_offset=topk_indices_offset,
seq_lens=seq_lens,
)
out_our = fast_kpool_topk_transform_fused(
score,
lengths,
pool_size,
topk,
page_table=page_table,
topk_indices_offset=topk_indices_offset,
seq_lens=seq_lens,
)
torch.cuda.synchronize()
out_ref = torch.sort(out_ref, dim=-1).values
out_our = torch.sort(out_our, dim=-1).values
torch.testing.assert_close(out_our, out_ref, atol=0, rtol=0)
if __name__ == "__main__":
sys.exit(pytest.main([__file__, "-v", "-s"]))