[JIT] Drop redundant per-kernel arch overrides (#32952)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
This commit is contained in:
DarkSharpness
2026-08-04 08:33:54 +08:00
committed by GitHub
co-authored by Claude Fable 5 Xiaoyu Zhang
parent 7eb27372b3
commit 1307968605
5 changed files with 123 additions and 76 deletions
+7 -10
View File
@@ -7,8 +7,8 @@ import subprocess
from tvm_ffi.libinfo import find_dlpack_include_path, find_include_path from tvm_ffi.libinfo import find_dlpack_include_path, find_include_path
from sglang.kernels.jit.utils import get_jit_cuda_arch, override_jit_cuda_arch from sglang.kernels.jit.utils import get_jit_cuda_arch
from sglang.kernels.jit.utils.arch import get_default_target_flags from sglang.kernels.jit.utils.arch import get_default_target_flags, make_jit_cuda_arch
from sglang.kernels.jit.utils.compile import DEFAULT_INCLUDE from sglang.kernels.jit.utils.compile import DEFAULT_INCLUDE
from sglang.kernels.jit.utils.deps import REGISTERED_DEPENDENCIES from sglang.kernels.jit.utils.deps import REGISTERED_DEPENDENCIES
@@ -69,27 +69,24 @@ def generate_clangd():
if args.cuda_target: if args.cuda_target:
assert args.cuda_target.count(".") == 1 assert args.cuda_target.count(".") == 1
major, minor = args.cuda_target.split(".") major, minor = args.cuda_target.split(".")
major, minor = int(major), int(minor) arch = make_jit_cuda_arch(int(major), int(minor))
context = override_jit_cuda_arch(major, minor)
context.__enter__()
else: else:
arch = get_jit_cuda_arch() arch = get_jit_cuda_arch()
major, minor = arch.major, f"{arch.minor}{arch.suffix}"
assert ( assert (
major > 0 arch.major > 0
), "Cannot detect CUDA architecture, please specify --cuda-target explicitly." ), "Cannot detect CUDA architecture, please specify --cuda-target explicitly."
compile_flags = [ compile_flags = [
"-xcuda", "-xcuda",
f"--cuda-gpu-arch=sm_{major}{minor}", f"--cuda-gpu-arch=sm_{arch.major}{arch.minor}{arch.suffix}",
"-Wall", "-Wall",
"-Wextra", "-Wextra",
*get_default_target_flags(), *get_default_target_flags(arch=arch),
*[f"-isystem{path}" for path in include_paths], *[f"-isystem{path}" for path in include_paths],
] ]
# NOTE: for local clangd (fix the missing cluster related macros) # NOTE: for local clangd (fix the missing cluster related macros)
if major >= 9: if arch.major >= 9:
compile_flags.append("-D_CG_LIMIT_INCLUDED_DEPENDENCIES=1") compile_flags.append("-D_CG_LIMIT_INCLUDED_DEPENDENCIES=1")
compile_flags.append("-D_CG_HAS_CLUSTER_GROUP=1") compile_flags.append("-D_CG_HAS_CLUSTER_GROUP=1")
+73 -11
View File
@@ -3,6 +3,10 @@
from __future__ import annotations from __future__ import annotations
import logging import logging
import os
import re
import shutil
import subprocess
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import dataclass from dataclasses import dataclass
from typing import List from typing import List
@@ -14,7 +18,6 @@ from sglang.kernels.jit.utils.common import (
is_hip_runtime, is_hip_runtime,
is_musa_runtime, is_musa_runtime,
) )
from sglang.srt.utils.common import get_cuda_version
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -34,18 +37,55 @@ class ArchInfo:
return f"-DSGL_CUDA_ARCH={self.major * 100 + self.minor * 10}" return f"-DSGL_CUDA_ARCH={self.major * 100 + self.minor * 10}"
@cache_once
def _jit_cuda_version() -> tuple[int, ...]:
"""CUDA version of the nvcc that JIT builds actually run.
The target has to match the compiler, not the toolkit PyTorch was built
against: a cu129 wheel on a CUDA 12.8 toolkit would otherwise select
`sm_120f`, which nvcc 12.8 rejects. Resolve nvcc the way tvm-ffi does
(`CUDA_HOME` / `CUDA_PATH`, then `$PATH`, then `/usr/local/cuda`) and fall
back to `torch.version.cuda` when it cannot be probed.
"""
cuda_home = os.environ.get("CUDA_HOME") or os.environ.get("CUDA_PATH")
if cuda_home is None:
nvcc_path = shutil.which("nvcc")
cuda_home = (
os.path.dirname(os.path.dirname(nvcc_path))
if nvcc_path is not None
else "/usr/local/cuda"
)
nvcc = os.path.join(cuda_home, "bin", "nvcc")
try:
output = subprocess.check_output([nvcc, "--version"], text=True)
match = re.search(r"release (\d+)\.(\d+)", output)
if match is not None:
return int(match.group(1)), int(match.group(2))
logger.warning("Cannot parse `%s --version` output: %s", nvcc, output)
except (OSError, subprocess.SubprocessError) as error:
logger.warning("Cannot run `%s --version`: %s", nvcc, error)
from sglang.srt.utils.common import get_cuda_version
return get_cuda_version()
def _cuda_arch_suffix(major: int, minor: int) -> str: def _cuda_arch_suffix(major: int, minor: int) -> str:
"""Mirror FlashInfer's `_normalize_cuda_arch`: 9.x/10.x+ -> "a"; 12.0 -> "f" """Mirror FlashInfer's `_normalize_cuda_arch`: 9.x/10.x+ -> "a"; 12.0 -> "f"
and 12.x (x>0) -> "a" (SM120/SM121 need separate cubins to avoid and 12.x (x>0) -> "a" (SM120/SM121 need separate cubins to avoid
cudaErrorIllegalInstruction, requires CUDA >= 12.9); below 9.0 -> plain. cudaErrorIllegalInstruction); below 9.0 -> plain.
Unlike FlashInfer, pre-12.9 CUDA falls back to plain instead of raising.
The family-specific "f" target needs a CUDA >= 12.9 nvcc; older toolkits
fall back to "a", which SM120 has had since 12.8. Plain sm_120 is never a
valid fallback: without an a/f target CUTLASS's SM120 atoms lose LDSM/STSM
and compile down to trap stubs (verified via SASS), which asserts at launch
instead of failing the build.
""" """
if major == 9: if major == 9:
return "a" return "a"
if major == 12: if major == 12:
if get_cuda_version() < (12, 9): if minor == 0 and _jit_cuda_version() >= (12, 9):
return "" return "f"
return "f" if minor == 0 else "a" return "a"
if major >= 10: if major >= 10:
return "a" return "a"
return "" return ""
@@ -71,7 +111,8 @@ def _init_jit_cuda_arch_once():
_CUDA_ARCH = ArchInfo(major, minor, suffix) _CUDA_ARCH = ArchInfo(major, minor, suffix)
def get_default_target_flags() -> List[str]: def get_default_target_flags(arch: ArchInfo | None = None) -> List[str]:
"""Default compile flags for `arch`, defaulting to the detected local GPU."""
if is_hip_runtime(): if is_hip_runtime():
flags = ["-DUSE_ROCM", "-std=c++20", "-O3"] flags = ["-DUSE_ROCM", "-std=c++20", "-O3"]
# Detect FP8 type based on GPU architecture # Detect FP8 type based on GPU architecture
@@ -86,20 +127,41 @@ def get_default_target_flags() -> List[str]:
flags.append("-DHIP_FP8_TYPE_E4M3=1") flags.append("-DHIP_FP8_TYPE_E4M3=1")
return flags return flags
else: else:
if arch is None:
arch = get_jit_cuda_arch()
return [ return [
get_jit_cuda_arch().jit_flag, arch.jit_flag,
"-std=c++20", "-std=c++20",
"-O3", "-O3",
"--expt-relaxed-constexpr", "--expt-relaxed-constexpr",
] ]
def make_jit_cuda_arch(major: int, minor: int) -> ArchInfo:
"""Build the JIT target for an explicitly requested capability."""
return ArchInfo(major, minor, _cuda_arch_suffix(major, minor))
@contextmanager @contextmanager
def override_jit_cuda_arch(major: int, minor: int, suffix: str = ""): def override_jit_cuda_arch(major: int, minor: int, suffix: str | None = None):
"""A context manager to temporarily override CUDA architecture.""" """A context manager to temporarily override CUDA architecture.
`suffix` defaults to the arch-specific target detection would pick for that
capability; pass it explicitly only to force a different one (an unsuffixed
target loses the arch-only instructions CUTLASS needs, see
`_cuda_arch_suffix`).
Kernels do not need this to reach an arch-specific target: `get_jit_cuda_arch`
already resolves the local GPU to its a/f target. Reach for it only to compile
for an arch the local GPU is not.
"""
global _CUDA_ARCH global _CUDA_ARCH
old_value = get_jit_cuda_arch() old_value = get_jit_cuda_arch()
_CUDA_ARCH = ArchInfo(major, minor, suffix) _CUDA_ARCH = (
make_jit_cuda_arch(major, minor)
if suffix is None
else ArchInfo(major, minor, suffix)
)
try: try:
yield yield
finally: finally:
@@ -15,7 +15,7 @@ from typing import TYPE_CHECKING
import torch import torch
from sglang.kernel_api_logging import debug_kernel_api from sglang.kernel_api_logging import debug_kernel_api
from sglang.kernels.jit.utils import cache_once, load_jit, override_jit_cuda_arch from sglang.kernels.jit.utils import cache_once, load_jit
if TYPE_CHECKING: if TYPE_CHECKING:
from tvm_ffi.module import Module from tvm_ffi.module import Module
@@ -29,22 +29,21 @@ ROPE_DIM = 64 # qk_rope_head_dim
def _jit_qprep_bf16_fp8_module() -> Module: def _jit_qprep_bf16_fp8_module() -> Module:
if torch.cuda.get_device_capability()[0] != 9: if torch.cuda.get_device_capability()[0] != 9:
raise RuntimeError("qprep_bf16_fp8_sm90 requires an SM90 (Hopper) GPU") raise RuntimeError("qprep_bf16_fp8_sm90 requires an SM90 (Hopper) GPU")
with override_jit_cuda_arch(9, 0, "a"): return load_jit(
return load_jit( "qprep_bf16_fp8_sm90",
"qprep_bf16_fp8_sm90", cuda_files=["qprep_bf16_fp8_sm90/entry.cuh"],
cuda_files=["qprep_bf16_fp8_sm90/entry.cuh"], cuda_wrappers=[("dispatch", "qprep_bf16_fp8_dispatch")],
cuda_wrappers=[("dispatch", "qprep_bf16_fp8_dispatch")], # Same minimal flag set as the sparse_mla_q8kv8_prefill_sm90 JIT
# Same minimal flag set as the sparse_mla_q8kv8_prefill_sm90 JIT # build (per-flag ablation there showed the rest are no-ops).
# build (per-flag ablation there showed the rest are no-ops). extra_cuda_cflags=[
extra_cuda_cflags=[ "-O3",
"-O3", "-DNDEBUG",
"-DNDEBUG", "-DCUTE_USE_PACKED_TUPLE=1",
"-DCUTE_USE_PACKED_TUPLE=1", "-DCUTLASS_ENABLE_TENSOR_CORE_MMA=1",
"-DCUTLASS_ENABLE_TENSOR_CORE_MMA=1", "--use_fast_math",
"--use_fast_math", ],
], extra_dependencies=["cutlass"],
extra_dependencies=["cutlass"], )
)
# torch._C._cuda_getCurrentRawStream returns the cudaStream_t pointer expected # torch._C._cuda_getCurrentRawStream returns the cudaStream_t pointer expected
@@ -11,7 +11,7 @@ from typing import TYPE_CHECKING, Optional
import torch import torch
from sglang.kernel_api_logging import debug_kernel_api from sglang.kernel_api_logging import debug_kernel_api
from sglang.kernels.jit.utils import cache_once, load_jit, override_jit_cuda_arch from sglang.kernels.jit.utils import cache_once, load_jit
from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.custom_op import register_custom_op
if TYPE_CHECKING: if TYPE_CHECKING:
@@ -59,20 +59,19 @@ def _q8kv8_cuda_flags() -> list[str]:
@cache_once @cache_once
def _jit_sparse_mla_q8kv8_prefill_module() -> Module: def _jit_sparse_mla_q8kv8_prefill_module() -> Module:
with override_jit_cuda_arch(9, 0, "a"): return load_jit(
return load_jit( "sparse_mla_q8kv8_prefill_sm90",
"sparse_mla_q8kv8_prefill_sm90", cuda_files=[
cuda_files=[ "sparse_mla_q8kv8_prefill_sm90/entry.cuh",
"sparse_mla_q8kv8_prefill_sm90/entry.cuh", ],
], cuda_wrappers=[
cuda_wrappers=[ ("dispatch", "sparse_prefill_q8kv8_dispatch"),
("dispatch", "sparse_prefill_q8kv8_dispatch"), ("dispatch_full", "sparse_prefill_q8kv8_dispatch_full"),
("dispatch_full", "sparse_prefill_q8kv8_dispatch_full"), ("dispatch_topk_length", "sparse_prefill_q8kv8_dispatch_topk_length"),
("dispatch_topk_length", "sparse_prefill_q8kv8_dispatch_topk_length"), ],
], extra_cuda_cflags=_q8kv8_cuda_flags(),
extra_cuda_cflags=_q8kv8_cuda_flags(), extra_dependencies=["cutlass"],
extra_dependencies=["cutlass"], )
)
# Pre-resolve entry-point callables on first use to avoid per-call module # Pre-resolve entry-point callables on first use to avoid per-call module
@@ -1,12 +1,11 @@
from __future__ import annotations from __future__ import annotations
from contextlib import contextmanager
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
import torch import torch
from sglang.kernel_api_logging import debug_kernel_api from sglang.kernel_api_logging import debug_kernel_api
from sglang.kernels.jit.utils import cache_once, load_jit, override_jit_cuda_arch from sglang.kernels.jit.utils import cache_once, load_jit
from sglang.srt.utils.common import is_sm120_supported from sglang.srt.utils.common import is_sm120_supported
from sglang.srt.utils.custom_op import register_custom_op from sglang.srt.utils.custom_op import register_custom_op
@@ -28,31 +27,22 @@ def _fp8_blockwise_cuda_flags() -> list[str]:
] ]
@contextmanager @cache_once
def _fp8_blockwise_arch_env(): def _jit_fp8_blockwise_module() -> Module:
"""Compile and cache the SM120 fp8 blockwise GEMM module (handles fp16 + bf16)."""
if not is_sm120_supported(): if not is_sm120_supported():
raise RuntimeError( raise RuntimeError(
"fp8_blockwise_scaled_mm JIT kernel requires SM120 (Blackwell)." "fp8_blockwise_scaled_mm JIT kernel requires SM120 (Blackwell)."
) )
major, minor = torch.cuda.get_device_capability() return load_jit(
# sm_*a target (e.g. sm_120a) required, not plain sm_120. "fp8_blockwise_scaled_mm",
with override_jit_cuda_arch(major, minor, suffix="a"): cuda_files=["gemm/fp8_blockwise/fp8_blockwise_scaled_mm_entry.cuh"],
yield cuda_wrappers=[
("fp8_blockwise_scaled_mm", "fp8_blockwise_scaled_mm"),
],
@cache_once extra_dependencies=["cutlass"],
def _jit_fp8_blockwise_module() -> Module: extra_cuda_cflags=_fp8_blockwise_cuda_flags(),
"""Compile and cache the SM120 fp8 blockwise GEMM module (handles fp16 + bf16).""" )
with _fp8_blockwise_arch_env():
return load_jit(
"fp8_blockwise_scaled_mm",
cuda_files=["gemm/fp8_blockwise/fp8_blockwise_scaled_mm_entry.cuh"],
cuda_wrappers=[
("fp8_blockwise_scaled_mm", "fp8_blockwise_scaled_mm"),
],
extra_dependencies=["cutlass"],
extra_cuda_cflags=_fp8_blockwise_cuda_flags(),
)
@register_custom_op( @register_custom_op(