diff --git a/python/sglang/jit_kernel/benchmark/bench_per_tensor_quant_fp8.py b/python/sglang/jit_kernel/benchmark/bench_per_tensor_quant_fp8.py index 8c19cb4b7..1fb0e45cb 100644 --- a/python/sglang/jit_kernel/benchmark/bench_per_tensor_quant_fp8.py +++ b/python/sglang/jit_kernel/benchmark/bench_per_tensor_quant_fp8.py @@ -1,10 +1,10 @@ -import os from typing import Optional, Tuple import torch import triton import triton.testing +from sglang.jit_kernel.benchmark.utils import is_in_ci from sglang.jit_kernel.per_tensor_quant_fp8 import per_tensor_quant_fp8 try: @@ -22,10 +22,7 @@ try: except ImportError: _is_hip = False -IS_CI = ( - os.getenv("CI", "false").lower() == "true" - or os.getenv("GITHUB_ACTIONS", "false").lower() == "true" -) +IS_CI = is_in_ci() fp8_type_ = torch.float8_e4m3fnuz if _is_hip else torch.float8_e4m3fn diff --git a/python/sglang/jit_kernel/benchmark/bench_qknorm.py b/python/sglang/jit_kernel/benchmark/bench_qknorm.py index 20ecd3b6e..7f5821e92 100644 --- a/python/sglang/jit_kernel/benchmark/bench_qknorm.py +++ b/python/sglang/jit_kernel/benchmark/bench_qknorm.py @@ -1,5 +1,4 @@ import itertools -import os from typing import Tuple import torch @@ -7,13 +6,11 @@ import triton import triton.testing from sgl_kernel import rmsnorm +from sglang.jit_kernel.benchmark.utils import is_in_ci from sglang.jit_kernel.norm import fused_inplace_qknorm from sglang.srt.utils import get_current_device_stream_fast -IS_CI = ( - os.getenv("CI", "false").lower() == "true" - or os.getenv("GITHUB_ACTIONS", "false").lower() == "true" -) +IS_CI = is_in_ci() alt_stream = torch.cuda.Stream() diff --git a/python/sglang/jit_kernel/benchmark/bench_rmsnorm.py b/python/sglang/jit_kernel/benchmark/bench_rmsnorm.py index 6d2a14824..1863c189f 100644 --- a/python/sglang/jit_kernel/benchmark/bench_rmsnorm.py +++ b/python/sglang/jit_kernel/benchmark/bench_rmsnorm.py @@ -1,5 +1,4 @@ import itertools -import os import torch import triton @@ -7,12 +6,10 @@ import triton.testing from flashinfer import rmsnorm as fi_rmsnorm from sgl_kernel import rmsnorm +from sglang.jit_kernel.benchmark.utils import is_in_ci from sglang.jit_kernel.norm import rmsnorm as jit_rmsnorm -IS_CI = ( - os.getenv("CI", "false").lower() == "true" - or os.getenv("GITHUB_ACTIONS", "false").lower() == "true" -) +IS_CI = is_in_ci() def sglang_aot_rmsnorm(