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
106 lines
3.6 KiB
Python
106 lines
3.6 KiB
Python
import sys
|
|
|
|
import pytest
|
|
import torch
|
|
|
|
from sglang.jit_kernel.mla_kv_pack_quantize_fp8 import mla_kv_pack_quantize_fp8
|
|
from sglang.jit_kernel.utils import get_ci_test_range
|
|
from sglang.test.ci.ci_register import register_amd_ci, register_cuda_ci
|
|
|
|
register_cuda_ci(est_time=60, stage="base-b-kernel-unit", runner_config="1-gpu-large")
|
|
register_amd_ci(est_time=60, suite="nightly-amd-kernel-1-gpu", nightly=True)
|
|
|
|
DEVICE = "cuda"
|
|
|
|
SHAPES = get_ci_test_range(
|
|
[(128, 64, 128), (64, 32, 64)],
|
|
[(128, 64, 128)],
|
|
)
|
|
NUM_HEADS = get_ci_test_range([8, 16, 32, 64], [16, 32])
|
|
BATCH_SIZES = get_ci_test_range(
|
|
[1, 4, 17, 64, 257, 1024, 4096, 16384],
|
|
[1, 64, 1024, 16384],
|
|
)
|
|
|
|
|
|
def _ref(k_nope, k_pe, v, k_scale_inv, v_scale_inv, fp8_dtype):
|
|
s, h, qk_nope = k_nope.shape
|
|
qk_rope = k_pe.shape[-1]
|
|
v_head = v.shape[-1]
|
|
if k_pe.dim() == 3:
|
|
k_pe = k_pe.squeeze(1)
|
|
|
|
k_bf16 = torch.empty(
|
|
(s, h, qk_nope + qk_rope), dtype=k_nope.dtype, device=k_nope.device
|
|
)
|
|
k_bf16[..., :qk_nope] = k_nope
|
|
k_bf16[..., qk_nope:] = k_pe.unsqueeze(1).expand(-1, h, -1)
|
|
|
|
k_fp8 = (k_bf16.float() * k_scale_inv).to(fp8_dtype)
|
|
v_fp8 = (v.float() * v_scale_inv).to(fp8_dtype)
|
|
return k_fp8, v_fp8
|
|
|
|
|
|
@pytest.mark.parametrize("dtype", [torch.bfloat16, torch.float16])
|
|
@pytest.mark.parametrize("shape", SHAPES)
|
|
@pytest.mark.parametrize("num_heads", NUM_HEADS)
|
|
@pytest.mark.parametrize("batch_size", BATCH_SIZES)
|
|
def test_correctness(dtype, shape, num_heads, batch_size):
|
|
qk_nope, qk_rope, v_head = shape
|
|
|
|
torch.manual_seed(0)
|
|
k_nope = torch.randn((batch_size, num_heads, qk_nope), dtype=dtype, device=DEVICE)
|
|
k_pe = torch.randn((batch_size, 1, qk_rope), dtype=dtype, device=DEVICE)
|
|
v = torch.randn((batch_size, num_heads, v_head), dtype=dtype, device=DEVICE)
|
|
|
|
k_scale_inv = 0.7
|
|
v_scale_inv = 1.3
|
|
|
|
k_fp8, v_fp8 = mla_kv_pack_quantize_fp8(
|
|
k_nope, k_pe, v, k_scale_inv=k_scale_inv, v_scale_inv=v_scale_inv
|
|
)
|
|
|
|
k_ref, v_ref = _ref(k_nope, k_pe, v, k_scale_inv, v_scale_inv, torch.float8_e4m3fn)
|
|
|
|
torch.testing.assert_close(k_fp8.float(), k_ref.float(), rtol=1e-2, atol=0.5)
|
|
torch.testing.assert_close(v_fp8.float(), v_ref.float(), rtol=1e-2, atol=0.5)
|
|
|
|
|
|
@pytest.mark.parametrize("dtype", [torch.bfloat16])
|
|
def test_strided_inputs(dtype):
|
|
s, h = 16, 32
|
|
qk_nope, qk_rope, v_head = 128, 64, 128
|
|
|
|
full = torch.randn(
|
|
(s, h, qk_nope * 2), dtype=dtype, device=DEVICE, requires_grad=False
|
|
)
|
|
k_nope = full[..., qk_nope:]
|
|
assert k_nope.stride(-1) == 1
|
|
|
|
k_pe = torch.randn((s, 1, qk_rope), dtype=dtype, device=DEVICE)
|
|
v = torch.randn((s, h, v_head), dtype=dtype, device=DEVICE)
|
|
|
|
k_fp8, v_fp8 = mla_kv_pack_quantize_fp8(k_nope, k_pe, v)
|
|
k_ref, v_ref = _ref(k_nope, k_pe, v, 1.0, 1.0, torch.float8_e4m3fn)
|
|
torch.testing.assert_close(k_fp8.float(), k_ref.float(), rtol=1e-2, atol=0.5)
|
|
torch.testing.assert_close(v_fp8.float(), v_ref.float(), rtol=1e-2, atol=0.5)
|
|
|
|
|
|
def test_kpe_2d_accepted():
|
|
s, h = 8, 16
|
|
qk_nope, qk_rope, v_head = 128, 64, 128
|
|
dtype = torch.bfloat16
|
|
|
|
k_nope = torch.randn((s, h, qk_nope), dtype=dtype, device=DEVICE)
|
|
k_pe = torch.randn((s, qk_rope), dtype=dtype, device=DEVICE)
|
|
v = torch.randn((s, h, v_head), dtype=dtype, device=DEVICE)
|
|
|
|
k_fp8, v_fp8 = mla_kv_pack_quantize_fp8(k_nope, k_pe, v)
|
|
k_ref, v_ref = _ref(k_nope, k_pe.unsqueeze(1), v, 1.0, 1.0, torch.float8_e4m3fn)
|
|
torch.testing.assert_close(k_fp8.float(), k_ref.float(), rtol=1e-2, atol=0.5)
|
|
torch.testing.assert_close(v_fp8.float(), v_ref.float(), rtol=1e-2, atol=0.5)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
sys.exit(pytest.main([__file__, "-v", "-s"]))
|