[sgl-kernel] fix b200 kernel ci (#13907)
Co-authored-by: HydraQYH <qyh820@outlook.com>
This commit is contained in:
@@ -271,36 +271,36 @@ jobs:
|
|||||||
|
|
||||||
echo "All benchmark tests completed!"
|
echo "All benchmark tests completed!"
|
||||||
|
|
||||||
# sgl-kernel-b200-test:
|
sgl-kernel-b200-test:
|
||||||
# needs: [check-changes, sgl-kernel-build-wheels]
|
needs: [check-changes, sgl-kernel-build-wheels]
|
||||||
# if: needs.check-changes.outputs.sgl_kernel == 'true'
|
if: needs.check-changes.outputs.sgl_kernel == 'true'
|
||||||
# runs-on: 4-gpu-b200
|
runs-on: 4-gpu-b200
|
||||||
# env:
|
env:
|
||||||
# RUNNER_LABELS: 4-gpu-b200
|
RUNNER_LABELS: 4-gpu-b200
|
||||||
# steps:
|
steps:
|
||||||
# - uses: actions/checkout@v4
|
- uses: actions/checkout@v4
|
||||||
|
|
||||||
# - name: Cleanup
|
- name: Cleanup
|
||||||
# run: |
|
run: |
|
||||||
# ls -alh sgl-kernel/dist || true
|
ls -alh sgl-kernel/dist || true
|
||||||
# rm -rf sgl-kernel/dist/* || true
|
rm -rf sgl-kernel/dist/* || true
|
||||||
|
|
||||||
# - name: Download artifacts
|
- name: Download artifacts
|
||||||
# uses: actions/download-artifact@v4
|
uses: actions/download-artifact@v4
|
||||||
# with:
|
with:
|
||||||
# path: sgl-kernel/dist/
|
path: sgl-kernel/dist/
|
||||||
# merge-multiple: true
|
merge-multiple: true
|
||||||
# pattern: wheel-python3.10-cuda12.9
|
pattern: wheel-python3.10-cuda12.9
|
||||||
|
|
||||||
# - name: Install dependencies
|
- name: Install dependencies
|
||||||
# run: |
|
run: |
|
||||||
# CUSTOM_BUILD_SGL_KERNEL=${{needs.check-changes.outputs.sgl_kernel}} IS_BLACKWELL=1 bash scripts/ci/ci_install_dependency.sh
|
CUSTOM_BUILD_SGL_KERNEL=${{needs.check-changes.outputs.sgl_kernel}} IS_BLACKWELL=1 bash scripts/ci/ci_install_dependency.sh
|
||||||
|
|
||||||
# - name: Run sgl-kernel unit tests on B200
|
- name: Run sgl-kernel unit tests on B200
|
||||||
# timeout-minutes: 30
|
timeout-minutes: 30
|
||||||
# run: |
|
run: |
|
||||||
# cd sgl-kernel
|
cd sgl-kernel
|
||||||
# pytest tests/
|
pytest tests/
|
||||||
|
|
||||||
# Adding a single CUDA13 smoke test to verify that the kernel builds and runs
|
# Adding a single CUDA13 smoke test to verify that the kernel builds and runs
|
||||||
# TODO: Add back this test when it can pass on CI
|
# TODO: Add back this test when it can pass on CI
|
||||||
@@ -1094,6 +1094,7 @@ jobs:
|
|||||||
sgl-kernel-unit-test,
|
sgl-kernel-unit-test,
|
||||||
sgl-kernel-mla-test,
|
sgl-kernel-mla-test,
|
||||||
sgl-kernel-benchmark-test,
|
sgl-kernel-benchmark-test,
|
||||||
|
sgl-kernel-b200-test,
|
||||||
|
|
||||||
multimodal-gen-test-1-gpu,
|
multimodal-gen-test-1-gpu,
|
||||||
multimodal-gen-test-2-gpu,
|
multimodal-gen-test-2-gpu,
|
||||||
|
|||||||
@@ -99,8 +99,8 @@ def is_sm90_supported(device=None) -> bool:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.skipif(
|
@pytest.mark.skipif(
|
||||||
not (is_sm100_supported() or is_sm90_supported()),
|
not is_sm90_supported(),
|
||||||
reason="fp8_blockwise_scaled_grouped_mm at sgl-kernel is only supported on sm100 or sm90",
|
reason="es_fp8_blockwise_scaled_grouped_mm at sgl-kernel is only supported on sm90",
|
||||||
)
|
)
|
||||||
@pytest.mark.parametrize("num_experts", [8, 16, 32, 64, 128])
|
@pytest.mark.parametrize("num_experts", [8, 16, 32, 64, 128])
|
||||||
@pytest.mark.parametrize("out_dtype", [torch.half, torch.bfloat16])
|
@pytest.mark.parametrize("out_dtype", [torch.half, torch.bfloat16])
|
||||||
|
|||||||
@@ -38,6 +38,12 @@ CAUSAL_TOPK = [(True, None), (False, None), (False, 128), (False, 2048)]
|
|||||||
DTYPE = [torch.float16, torch.bfloat16]
|
DTYPE = [torch.float16, torch.bfloat16]
|
||||||
|
|
||||||
|
|
||||||
|
def is_sm90_supported(device=None) -> bool:
|
||||||
|
return (torch.cuda.get_device_capability(device)[0] == 9) and (
|
||||||
|
torch.version.cuda >= "12.3"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def quantize_k_cache(
|
def quantize_k_cache(
|
||||||
input_k_cache: torch.Tensor, # (num_blocks, block_size, h_k, d)
|
input_k_cache: torch.Tensor, # (num_blocks, block_size, h_k, d)
|
||||||
dv: int,
|
dv: int,
|
||||||
@@ -362,6 +368,7 @@ def test_flashmla_prefill(
|
|||||||
torch.testing.assert_close(ans_lse, ref_lse, atol=1e-6, rtol=2.01 / 65536)
|
torch.testing.assert_close(ans_lse, ref_lse, atol=1e-6, rtol=2.01 / 65536)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not is_sm90_supported(), reason="SM90 required for FP8 support")
|
||||||
@pytest.mark.parametrize("b", B_DECODE)
|
@pytest.mark.parametrize("b", B_DECODE)
|
||||||
@pytest.mark.parametrize("s_q", S_Q_DECODE)
|
@pytest.mark.parametrize("s_q", S_Q_DECODE)
|
||||||
@pytest.mark.parametrize("s_k", S_K_DECODE)
|
@pytest.mark.parametrize("s_k", S_K_DECODE)
|
||||||
@@ -512,6 +519,7 @@ def test_flash_mla_decode(
|
|||||||
torch.testing.assert_close(lse_ans, lse_ref, atol=1e-6, rtol=8.01 / 65536)
|
torch.testing.assert_close(lse_ans, lse_ref, atol=1e-6, rtol=8.01 / 65536)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.skipif(not is_sm90_supported(), reason="SM90 required for FP8 support")
|
||||||
@pytest.mark.parametrize("b", [128])
|
@pytest.mark.parametrize("b", [128])
|
||||||
@pytest.mark.parametrize("s_q", [1, 2])
|
@pytest.mark.parametrize("s_q", [1, 2])
|
||||||
@pytest.mark.parametrize("mean_sk", [4096, 8192, 16384])
|
@pytest.mark.parametrize("mean_sk", [4096, 8192, 16384])
|
||||||
|
|||||||
Reference in New Issue
Block a user