项目文件夹

文件
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

246 行
7.0 KiB
Python

import unittest
import sgl_kernel # noqa: F401
import torch
import torch.nn.functional as F
from utils import parametrize, 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")
flash_attn_varlen_func = torch.ops.sgl_kernel.flash_attn_varlen_func
torch.manual_seed(1234)
def flash_attn_varlen_ref(
q,
k,
v,
cu_seqlens_q,
cu_seqlens_k,
is_causal,
enable_gqa,
):
cu_q = cu_seqlens_q.tolist()
cu_k = cu_seqlens_k.tolist()
batch = len(cu_k) - 1
# [T, H, D] -> [1, H, T, D]
q, k, v = [x.unsqueeze(0).transpose(1, 2) for x in [q, k, v]]
B, H, T, D = q.shape
out = torch.empty(B, H, T, v.size(-1), dtype=q.dtype)
for b in range(batch):
start_q, end_q = cu_q[b], cu_q[b + 1]
start_k, end_k = cu_k[b], cu_k[b + 1]
out[:, :, start_q:end_q, :] = F.scaled_dot_product_attention(
q[:, :, start_q:end_q, :],
k[:, :, start_k:end_k, :],
v[:, :, start_k:end_k, :],
is_causal=is_causal,
enable_gqa=enable_gqa,
)
# [1, H, T, D] -> [T, H, D]
return out.transpose(1, 2).squeeze(0)
# faster version ref kernel for non varlen case
def flash_attn_non_varlen_ref(
q,
k,
v,
cu_seqlens_q,
cu_seqlens_k,
is_causal,
enable_gqa,
):
cu_q = cu_seqlens_q.tolist()
cu_k = cu_seqlens_k.tolist()
batch = len(cu_k) - 1
B_T, H, D = q.shape
T = B_T // batch
# [T, H, D] -> [1, H, T, D]
q, k, v = [x.reshape(batch, T, H, D).transpose(1, 2) for x in [q, k, v]]
out = F.scaled_dot_product_attention(
q,
k,
v,
is_causal=is_causal,
enable_gqa=enable_gqa,
)
# [B, H, T, D] -> [B * T, H, D]
return out.transpose(1, 2).reshape(batch * T, H, D)
class TestFlashAttn(CustomTestCase):
@parametrize(
batch=[4],
max_seqlen_q=[35, 96],
max_seqlen_k=[35, 96],
num_heads=[16],
num_heads_kv=[16, 2],
head_dim=[32, 48], # test when D is not 32x
head_dim_v=[32],
is_causal=[True, False],
)
def test_flash_attn_varlen(
self,
batch,
max_seqlen_q,
max_seqlen_k,
num_heads,
num_heads_kv,
head_dim,
head_dim_v,
is_causal,
):
dtype = torch.bfloat16
# random seqlens for k and kv
seqlens_q = torch.randint(1, max_seqlen_q, (batch,), dtype=torch.int32)
seqlens_k = torch.randint(1, max_seqlen_k, (batch,), dtype=torch.int32)
cu_seqlens_q = torch.zeros((batch + 1,), dtype=torch.int32)
cu_seqlens_k = torch.zeros((batch + 1,), dtype=torch.int32)
cu_seqlens_q[1:] = torch.cumsum(seqlens_q, 0)
cu_seqlens_k[1:] = torch.cumsum(seqlens_k, 0)
sum_seqlen_q = seqlens_q.sum().item()
sum_seqlen_k = seqlens_k.sum().item()
q = torch.randn(sum_seqlen_q, num_heads, head_dim).to(dtype)
k = torch.randn(sum_seqlen_k, num_heads_kv, head_dim).to(dtype)
v = torch.randn(sum_seqlen_k, num_heads_kv, head_dim_v).to(dtype)
out_ref = flash_attn_varlen_ref(
q,
k,
v,
cu_seqlens_q,
cu_seqlens_k,
is_causal=is_causal,
enable_gqa=num_heads != num_heads_kv,
)
out = flash_attn_varlen_func(
q,
k,
v,
cu_seqlens_q,
cu_seqlens_k,
seqlens_q.max().item(),
seqlens_k.max().item(),
is_causal,
)
atol = rtol = precision[dtype]
torch.testing.assert_close(out_ref, out, atol=atol, rtol=rtol)
# test with large size to capture overflow issue
@parametrize(
batch=[4097],
max_seqlen_q=[4097],
max_seqlen_k=[4097],
num_heads=[4],
num_heads_kv=[4],
head_dim=[32],
head_dim_v=[32],
is_causal=[False],
)
def test_flash_attn_large_size(
self,
batch,
max_seqlen_q,
max_seqlen_k,
num_heads,
num_heads_kv,
head_dim,
head_dim_v,
is_causal,
):
dtype = torch.bfloat16
# test the non varlen case
seqlens_q = torch.full((batch,), max_seqlen_q, dtype=torch.int32)
seqlens_k = torch.full((batch,), max_seqlen_k, dtype=torch.int32)
cu_seqlens_q = torch.zeros((batch + 1,), dtype=torch.int32)
cu_seqlens_k = torch.zeros((batch + 1,), dtype=torch.int32)
cu_seqlens_q[1:] = torch.cumsum(seqlens_q, 0)
cu_seqlens_k[1:] = torch.cumsum(seqlens_k, 0)
sum_seqlen_q = seqlens_q.sum().item()
sum_seqlen_k = seqlens_k.sum().item()
q = torch.randn(sum_seqlen_q, num_heads, head_dim).to(dtype)
k = torch.randn(sum_seqlen_k, num_heads_kv, head_dim).to(dtype)
v = torch.randn(sum_seqlen_k, num_heads_kv, head_dim_v).to(dtype)
out_ref = flash_attn_non_varlen_ref(
q,
k,
v,
cu_seqlens_q,
cu_seqlens_k,
is_causal=is_causal,
enable_gqa=num_heads != num_heads_kv,
)
out = flash_attn_varlen_func(
q,
k,
v,
cu_seqlens_q,
cu_seqlens_k,
seqlens_q.max().item(),
seqlens_k.max().item(),
is_causal,
)
atol = rtol = precision[dtype]
torch.testing.assert_close(out_ref, out, atol=atol, rtol=rtol)
def _test_flash_attn_large_seq_causal_mask_once(self, seqlens):
dtype = torch.bfloat16
num_heads = 8
num_heads_kv = 2
head_dim = 64
seqlens_t = torch.tensor(seqlens, dtype=torch.int32)
cu_seqlens = torch.zeros(len(seqlens) + 1, dtype=torch.int32)
cu_seqlens[1:] = torch.cumsum(seqlens_t, 0)
total = cu_seqlens[-1].item()
max_seqlen = seqlens_t.max().item()
q = torch.randn(total, num_heads, head_dim, dtype=dtype)
k = torch.randn(total, num_heads_kv, head_dim, dtype=dtype)
v = torch.randn(total, num_heads_kv, head_dim, dtype=dtype)
out_ref = flash_attn_varlen_ref(
q, k, v, cu_seqlens, cu_seqlens, is_causal=True, enable_gqa=True
)
out = flash_attn_varlen_func(
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, True
)
atol = rtol = precision[dtype]
torch.testing.assert_close(out_ref, out, atol=atol, rtol=rtol)
def test_flash_attn_large_seq_causal_mask(self):
# Non-varlen path: single sequence, has_varlen_sequences returns False
# → dispatches to flash_attn_kernel_impl.
self._test_flash_attn_large_seq_causal_mask_once([5000])
# Varlen path: sequences with different lengths, has_varlen_sequences
# returns True → dispatches to flash_attn_varlen_kernel_impl
self._test_flash_attn_large_seq_causal_mask_once([5000, 4999])
if __name__ == "__main__":
unittest.main()