sgl-project--sglang
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
163 行
4.9 KiB
Python
163 行
4.9 KiB
Python
import unittest
|
|
|
|
import torch
|
|
from torch.nn.functional import scaled_dot_product_attention
|
|
from utils import precision
|
|
|
|
from sglang.test.ci.ci_register import register_cpu_ci
|
|
from sglang.test.test_utils import CustomTestCase
|
|
|
|
register_cpu_ci(est_time=10, suite="base-b-test-cpu")
|
|
|
|
torch.manual_seed(1234)
|
|
|
|
|
|
class TestMLA(CustomTestCase):
|
|
def _run_sdpa_forward_decode(
|
|
self,
|
|
query: torch.Tensor,
|
|
output: torch.Tensor,
|
|
k_cache: torch.Tensor,
|
|
v_cache: torch.Tensor,
|
|
key: torch.Tensor,
|
|
loc: torch.Tensor,
|
|
req_to_token: torch.Tensor,
|
|
req_pool_indices: torch.Tensor,
|
|
seq_lens: torch.Tensor,
|
|
scaling=None,
|
|
enable_gqa=False,
|
|
causal=False,
|
|
):
|
|
# set kv cache
|
|
k_cache[loc] = key
|
|
|
|
# [num_tokens, num_heads, head_size] -> [num_heads, num_tokens, head_size]
|
|
query = query.movedim(0, query.dim() - 2)
|
|
|
|
start_q, start_kv = 0, 0
|
|
for seq_idx in range(seq_lens.shape[0]):
|
|
seq_len_q = 1
|
|
seq_len_kv = seq_lens[seq_idx]
|
|
end_q = start_q + seq_len_q
|
|
end_kv = start_kv + seq_len_kv
|
|
|
|
per_req_query = query[:, start_q:end_q, :]
|
|
|
|
# get key and value from cache. per_req_tokens contains the kv cache
|
|
# index for each token in the sequence.
|
|
req_pool_idx = req_pool_indices[seq_idx]
|
|
per_req_tokens = req_to_token[req_pool_idx, :seq_len_kv]
|
|
per_req_key = k_cache[per_req_tokens].movedim(0, query.dim() - 2)
|
|
per_req_value = v_cache[per_req_tokens].movedim(0, query.dim() - 2)
|
|
|
|
per_req_out = (
|
|
scaled_dot_product_attention(
|
|
per_req_query.unsqueeze(0),
|
|
per_req_key.unsqueeze(0),
|
|
per_req_value.unsqueeze(0),
|
|
enable_gqa=enable_gqa,
|
|
scale=scaling,
|
|
is_causal=causal,
|
|
)
|
|
.squeeze(0)
|
|
.movedim(query.dim() - 2, 0)
|
|
)
|
|
output[start_q:end_q, :, :] = per_req_out
|
|
start_q, start_kv = end_q, end_kv
|
|
|
|
return output
|
|
|
|
def _test_grouped_decode_attention_once(self, B, H_Q, H_KV, D, D_V, seq_len):
|
|
dtype = torch.bfloat16
|
|
|
|
total_tokens = B * seq_len
|
|
sm_scale = 1.0 / (D**0.5)
|
|
logit_cap = 0.0
|
|
num_kv_splits = 8
|
|
enable_gqa = H_Q != H_KV
|
|
|
|
# q represents the new token being generated, one per batch
|
|
q = torch.randn(B, H_Q, D, dtype=dtype)
|
|
|
|
# k_buffer and v_buffer represent all previous tokens
|
|
k_buffer = torch.randn(total_tokens, H_KV, D, dtype=dtype)
|
|
v_buffer = k_buffer.narrow(2, 0, D_V)
|
|
|
|
key = torch.randn(B, H_KV, D, dtype=dtype)
|
|
value = key.narrow(2, 0, D_V)
|
|
# make sure no duplicates in loc
|
|
loc = torch.randperm(total_tokens)[:B].to(torch.int64)
|
|
|
|
k_buffer2 = k_buffer.clone()
|
|
v_buffer2 = k_buffer2.narrow(2, 0, D_V)
|
|
|
|
# o will have the same shape as q
|
|
o = torch.zeros(B, H_Q, D_V, dtype=dtype)
|
|
o_grouped = torch.zeros(B, H_Q, D_V, dtype=dtype)
|
|
|
|
req_to_token = torch.arange(total_tokens).reshape(B, seq_len).to(torch.int32)
|
|
b_req_idx = torch.arange(B).to(torch.int64)
|
|
b_seq_len = torch.full((B,), seq_len).to(torch.int64)
|
|
|
|
attn_logits = torch.empty(
|
|
(B, H_Q, num_kv_splits, D_V + 1),
|
|
dtype=torch.float32,
|
|
)
|
|
|
|
torch.ops.sgl_kernel.decode_attention_cpu(
|
|
q,
|
|
k_buffer2,
|
|
v_buffer2,
|
|
o,
|
|
key,
|
|
value,
|
|
loc,
|
|
attn_logits,
|
|
req_to_token,
|
|
b_req_idx,
|
|
b_seq_len,
|
|
sm_scale,
|
|
logit_cap,
|
|
False,
|
|
0,
|
|
None,
|
|
None,
|
|
)
|
|
|
|
self._run_sdpa_forward_decode(
|
|
q,
|
|
o_grouped,
|
|
k_buffer,
|
|
v_buffer,
|
|
key,
|
|
loc,
|
|
req_to_token,
|
|
b_req_idx,
|
|
b_seq_len,
|
|
scaling=sm_scale,
|
|
enable_gqa=enable_gqa,
|
|
)
|
|
|
|
cos_sim = torch.nn.functional.cosine_similarity(
|
|
o.flatten(), o_grouped.flatten(), dim=0
|
|
)
|
|
atol = rtol = precision[q.dtype]
|
|
self.assertGreater(cos_sim.item(), 0.99)
|
|
torch.testing.assert_close(o, o_grouped, atol=atol, rtol=rtol)
|
|
torch.testing.assert_close(k_buffer, k_buffer2, atol=atol, rtol=rtol)
|
|
torch.testing.assert_close(v_buffer, v_buffer2, atol=atol, rtol=rtol)
|
|
|
|
def test_grouped_decode_attention(self):
|
|
configs = [
|
|
(1, 22, 1, 576, 512, 8 * 111),
|
|
(4, 22, 1, 576, 512, 8 * 128),
|
|
(40, 22, 1, 576, 512, 8 * 133),
|
|
]
|
|
|
|
for B, H_Q, H_KV, D, D_V, seqlen in configs:
|
|
self._test_grouped_decode_attention_once(B, H_Q, H_KV, D, D_V, seqlen)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|