[Apple Silicon] Add Metal kernel support in sgl-kernel (#23449)

Signed-off-by: Xiaodong Ye <yeahdongcn@gmail.com>
This commit is contained in:
R0CKSTAR
2026-05-11 17:54:27 -07:00
committed by GitHub
parent de098f4f4d
commit 74d70af09a
10 changed files with 585 additions and 202 deletions
@@ -6,6 +6,19 @@ metatags:
This document describes how run SGLang on Apple Silicon using [Metal (MLX)](https://opensource.apple.com/projects/mlx/). If you encounter issues or have questions, please [open an issue](https://github.com/sgl-project/sglang/issues).
## Prerequisites
Building the native Metal kernels in `sgl-kernel` requires the Apple
toolchain (`clang++`, the Metal framework headers, and `xcrun`). These ship
with the **Xcode Command Line Tools**, which cannot be installed via `pip`:
```bash
xcode-select --install
```
If you have the full Xcode app installed, the Command Line Tools are already
available. You can verify with `xcode-select -p && xcrun --find metal`.
## Install SGLang
You can install SGLang using one of the methods below.
@@ -17,8 +30,15 @@ You can install SGLang using one of the methods below.
git clone https://github.com/sgl-project/sglang.git
cd sglang
# Install sglang python package
pip install --upgrade pip
# Create and activate a virtual environment
uv venv -p 3.12 sglang-metal
source sglang-metal/bin/activate
# (Optional) Compile sgl-kernel
uv pip install --upgrade pip
uv run sgl-kernel/setup_metal.py install
# Install sglang python package along with diffusion support
rm -f python/pyproject.toml && mv python/pyproject_other.toml python/pyproject.toml
uv pip install -e "python[all_mps]"
```
@@ -22,7 +22,7 @@ pip install --upgrade pip
cd sgl-kernel
python setup_musa.py install
# Install sglang python package
# Install sglang python package along with diffusion support
cd ..
rm -f python/pyproject.toml && mv python/pyproject_other.toml python/pyproject.toml
pip install -e "python[all_musa]"
@@ -120,7 +120,7 @@ git clone https://github.com/sgl-project/sglang.git
cd sglang
# Create and activate a virtual environment
uv venv -p 3.11 sglang-diffusion
uv venv -p 3.12 sglang-diffusion
source sglang-diffusion/bin/activate
# Install the Python packages
+3 -1
View File
@@ -3,9 +3,10 @@
# Install stubs early for platforms where certain dependencies are unavailable
# (e.g. macOS/MPS has no triton, and torch.mps lacks Stream / set_device /
# get_device_properties). This must run before any downstream imports.
import platform as _platform
import sys as _sys
if _sys.platform == "darwin":
if _sys.platform == "darwin" and _platform.machine() == "arm64":
try:
import torch as _torch
@@ -22,6 +23,7 @@ if _sys.platform == "darwin":
del _torch
except ImportError:
pass
del _platform
del _sys
from sglang.srt.utils.hf_transformers_patches import apply_all as _apply_hf_patches
+21
View File
@@ -0,0 +1,21 @@
# sgl-kernel Metal kernels
Custom Apple Metal kernels for the MLX backend on Apple Silicon. Shader sources (`*.metal`) and C++ host / nanobind sources (`*.cpp`) in this directory are compiled by [`sgl-kernel/setup_metal.py`](../../setup_metal.py) into the `sgl_kernel._metal` extension and the `sgl_metal_kernels.metallib` archive, and exposed through Python wrappers in [`python/sgl_kernel/metal.py`](../../python/sgl_kernel/metal.py).
## Kernels
| Kernel | Description | Tested on |
| --- | --- | --- |
| _none yet_ | — | — |
## Adding a new Metal kernel
1. Add the shader under `csrc/metal/<kernel>.metal`.
2. Add the C++ host / nanobind binding under `csrc/metal/<kernel>.cpp`, exporting the entry point on the `sgl_kernel._metal` module.
3. Append both files to `metal_shader_sources` and `cxx_sources` in [`sgl-kernel/setup_metal.py`](../../setup_metal.py).
4. Add a Python wrapper in [`python/sgl_kernel/metal.py`](../../python/sgl_kernel/metal.py) that validates input shapes/dtypes and calls `mx.eval` on its operands before invoking the AOT C++ entry point.
5. Add a test under [`sgl-kernel/tests/`](../../tests) and update the **Kernels** table above with a short description and the hardware / OS / MLX version the kernel was validated on.
## Note on `placeholder.metal` / `placeholder.cpp`
`placeholder.metal` and `placeholder.cpp` are intentionally empty. They exist only so that `setup_metal.py` has at least one shader source and one C++ source to compile, allowing the `sgl_kernel._metal` extension and the `sgl_metal_kernels.metallib` archive to build successfully before any real Metal kernels have been added. Both files (and their entries in `metal_shader_sources` / `cxx_sources` in `setup_metal.py`) MUST be removed once the first real kernel lands.
+60 -48
View File
@@ -1,23 +1,38 @@
import torch
from sgl_kernel.debug_utils import maybe_wrap_debug_kernel
from sgl_kernel.load_utils import _load_architecture_specific_ops, _preload_cuda_library
import platform
import sys
# Initialize the ops library based on current GPU
common_ops = _load_architecture_specific_ops()
from sgl_kernel.version import __version__ # noqa: F401
# Preload the CUDA library to avoid the issue of libcudart.so.12 not found
if torch.version.cuda is not None:
# On macOS only the Metal extension is shipped; skip CUDA op loading and
# re-exports so those symbols are not exposed on Apple Silicon.
if sys.platform == "darwin" and platform.machine() == "arm64":
from sgl_kernel.metal import *
else:
import torch
from sgl_kernel.debug_utils import maybe_wrap_debug_kernel
from sgl_kernel.load_utils import (
_load_architecture_specific_ops,
_preload_cuda_library,
)
# Initialize the ops library based on current GPU
common_ops = _load_architecture_specific_ops()
# Preload the CUDA library to avoid the issue of libcudart.so.12 not found
if torch.version.cuda is not None:
_preload_cuda_library()
from sgl_kernel.allreduce import *
from sgl_kernel.attention import (
from sgl_kernel.allreduce import *
from sgl_kernel.attention import (
cutlass_mla_decode,
cutlass_mla_get_workspace_size,
merge_state_v2,
)
from sgl_kernel.cutlass_moe import cutlass_w4a8_moe_mm, get_cutlass_w4a8_moe_mm_data
from sgl_kernel.elementwise import (
)
from sgl_kernel.cutlass_moe import (
cutlass_w4a8_moe_mm,
get_cutlass_w4a8_moe_mm_data,
)
from sgl_kernel.elementwise import (
concat_mla_absorb_q,
concat_mla_k,
copy_to_gpu_no_ce,
@@ -29,13 +44,13 @@ from sgl_kernel.elementwise import (
rmsnorm,
rotary_embedding,
silu_and_mul,
)
from sgl_kernel.expert_specialization import (
)
from sgl_kernel.expert_specialization import (
es_fp8_blockwise_scaled_grouped_mm,
es_sm100_mxfp8_blockscaled_grouped_mm,
es_sm100_mxfp8_blockscaled_grouped_quant,
)
from sgl_kernel.gemm import (
)
from sgl_kernel.gemm import (
awq_dequantize,
bmm_fp8,
dsv3_fused_a_gemm,
@@ -52,23 +67,23 @@ from sgl_kernel.gemm import (
sgl_per_token_group_quant_int8,
sgl_per_token_quant_fp8,
shuffle_rows,
)
from sgl_kernel.grammar import apply_token_bitmask_inplace_cuda
from sgl_kernel.kvcacheio import (
)
from sgl_kernel.grammar import apply_token_bitmask_inplace_cuda
from sgl_kernel.kvcacheio import (
transfer_kv_all_layer,
transfer_kv_all_layer_mla,
transfer_kv_per_layer,
transfer_kv_per_layer_mla,
)
from sgl_kernel.mamba import (
)
from sgl_kernel.mamba import (
causal_conv1d_fn_cpu,
causal_conv1d_fwd,
causal_conv1d_update,
causal_conv1d_update_cpu,
chunk_gated_delta_rule_cpu,
)
from sgl_kernel.memory import weak_ref_tensor
from sgl_kernel.moe import (
)
from sgl_kernel.memory import weak_ref_tensor
from sgl_kernel.moe import (
apply_shuffle_mul_sum,
fp8_blockwise_scaled_grouped_mm,
fused_qk_norm_rope,
@@ -80,38 +95,38 @@ from sgl_kernel.moe import (
prepare_moe_input,
topk_sigmoid,
topk_softmax,
)
from sgl_kernel.quantization import (
)
from sgl_kernel.quantization import (
ggml_dequantize,
ggml_moe_a8,
ggml_moe_a8_vec,
ggml_moe_get_block_size,
ggml_mul_mat_a8,
ggml_mul_mat_vec_a8,
)
from sgl_kernel.sampling import (
)
from sgl_kernel.sampling import (
top_k_renorm_prob,
top_p_renorm_prob,
)
from sgl_kernel.speculative import (
)
from sgl_kernel.speculative import (
build_tree_kernel_efficient,
reconstruct_indices_from_tree_mask,
segment_packbits,
tree_speculative_sampling_target_only,
verify_tree_greedy,
)
from sgl_kernel.top_k import (
)
from sgl_kernel.top_k import (
fast_topk,
fast_topk_transform_fused,
fast_topk_transform_ragged_fused,
fast_topk_v2,
)
from sgl_kernel.version import __version__
)
from sgl_kernel.version import __version__
if torch.version.hip is not None:
if torch.version.hip is not None:
from sgl_kernel.elementwise import gelu_quick
if hasattr(torch.version, "musa") and torch.version.musa is not None:
if hasattr(torch.version, "musa") and torch.version.musa is not None:
from sgl_kernel.musa import (
musa_batched_rotary_embedding_contiguous,
musa_fused_gemv,
@@ -120,8 +135,7 @@ if hasattr(torch.version, "musa") and torch.version.musa is not None:
musa_rotary_embedding_contiguous,
)
_DEBUG_EXPORT_NAMES = [
_DEBUG_EXPORT_NAMES = [
"apply_shuffle_mul_sum",
"apply_token_bitmask_inplace_cuda",
"awq_dequantize",
@@ -185,28 +199,26 @@ _DEBUG_EXPORT_NAMES = [
"tree_speculative_sampling_target_only",
"verify_tree_greedy",
"weak_ref_tensor",
]
]
if torch.version.hip is not None:
if torch.version.hip is not None:
_DEBUG_EXPORT_NAMES.append("gelu_quick")
for _name in _DEBUG_EXPORT_NAMES:
for _name in _DEBUG_EXPORT_NAMES:
if _name in globals():
globals()[_name] = maybe_wrap_debug_kernel(
globals()[_name], f"sgl_kernel.{_name}"
)
del _name
del _DEBUG_EXPORT_NAMES
del _name
del _DEBUG_EXPORT_NAMES
def create_greenctx_stream_by_value(*args, **kwargs):
def create_greenctx_stream_by_value(*args, **kwargs):
from sgl_kernel.spatial import create_greenctx_stream_by_value as _impl
return _impl(*args, **kwargs)
def get_sm_available(*args, **kwargs):
def get_sm_available(*args, **kwargs):
from sgl_kernel.spatial import get_sm_available as _impl
return _impl(*args, **kwargs)
+30
View File
@@ -0,0 +1,30 @@
"""Python entry points for the sgl_kernel Metal extension."""
from __future__ import annotations
from pathlib import Path
from typing import TYPE_CHECKING
if TYPE_CHECKING:
import mlx.core as mx
_METALLIB_NAME = "sgl_metal_kernels.metallib"
try:
from . import _metal
_metallib_path = Path(_metal.__file__).resolve().parent / _METALLIB_NAME
if not _metallib_path.is_file():
raise ImportError(
f"{_METALLIB_NAME} not found next to sgl_kernel._metal at {_metallib_path}"
)
_metal.register_library(str(_metallib_path))
except ImportError as _exc: # pragma: no cover - import guarded at call time
_metal = None
_IMPORT_ERROR: Exception | None = _exc
else:
_IMPORT_ERROR = None
# Python wrappers for the compiled `_metal.*` entry points go below. Each
# wrapper validates input shapes/dtypes and calls `mx.eval` on its operands
# before invoking the AOT C++ entry point.
+298
View File
@@ -0,0 +1,298 @@
# Copyright 2026 SGLang Team. All Rights Reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
import importlib
import os
import platform
import shutil
import subprocess
import sys
import sysconfig
from pathlib import Path
root = Path(__file__).parent.resolve()
_BUILD_REQUIRES = [
("setuptools", "setuptools"),
("mlx", "mlx"),
("nanobind", "nanobind"),
]
def _ensure_toolchain():
if sys.platform != "darwin" or platform.machine() != "arm64":
raise SystemExit("setup_metal.py only supports macOS (Apple Silicon).")
if shutil.which("c++") is None or shutil.which("xcrun") is None:
raise SystemExit(
"Apple toolchain not found. Install the Xcode Command Line Tools "
"with `xcode-select --install` (or a full Xcode install) and retry."
)
try:
subprocess.check_output(
["xcrun", "-sdk", "macosx", "metal", "--version"],
stderr=subprocess.STDOUT,
)
except (subprocess.CalledProcessError, FileNotFoundError) as exc:
raise SystemExit(
"Apple Metal shader compiler not found. Install a full Xcode "
"(not just Command Line Tools) so that `xcrun -sdk macosx metal` "
"is available, then retry."
) from exc
def _ensure_build_requires():
missing = []
for import_name, pip_name in _BUILD_REQUIRES:
try:
importlib.import_module(import_name)
except ImportError:
missing.append(pip_name)
if not missing:
return
print(
f"[sgl-kernel:metal] installing build requirements: {missing}",
flush=True,
)
subprocess.check_call(
[sys.executable, "-m", "pip", "install", "--upgrade", *missing]
)
# Section 1: Prerequisites
_ensure_toolchain()
_ensure_build_requires()
os.chdir(root)
# Section 2: Build and install
from setuptools import Extension, find_packages, setup # noqa: E402
from setuptools.command.build_ext import build_ext # noqa: E402
def _get_version():
with open(root / "pyproject.toml") as f:
for line in f:
if line.startswith("version"):
return line.split("=")[1].strip().strip('"')
operator_namespace = "sgl_kernel"
metallib_name = "sgl_metal_kernels.metallib"
# Metal shader sources (compiled with `xcrun metal`) and C++ host sources
# (compiled with `c++`). Add new kernels by appending to these lists.
metal_shader_sources = [
"csrc/metal/placeholder.metal",
]
cxx_sources = [
"csrc/metal/placeholder.cpp",
]
# Header search paths shared by both the Metal shader compiler and the C++
# host compiler.
include_dirs = [
root / "csrc",
root / "csrc" / "metal",
]
cxx_flags = ["-std=c++17", "-O3", "-fvisibility=hidden"]
metal_flags = ["-O3"]
frameworks = ["Metal", "Foundation", "QuartzCore"]
libraries = ["mlx"]
class BuildMetalExtension(build_ext):
def build_extension(self, ext):
if sys.platform != "darwin" or platform.machine() != "arm64":
raise RuntimeError("setup_metal.py only supports macOS (Apple Silicon).")
ext_path = Path(self.get_ext_fullpath(ext.name))
ext_path.parent.mkdir(parents=True, exist_ok=True)
# Use ccache for the C++ compiles when it is on PATH. ccache does not
# support `.metal` sources (unsupported source language), so the Metal
# shader compile is left untouched.
ccache = shutil.which("ccache")
cxx_cmd = [ccache, "c++"] if ccache else ["c++"]
if ccache:
print(f"[sgl-kernel:metal] using ccache at {ccache}", flush=True)
python_exe = Path(sys.executable)
python_include = Path(sysconfig.get_paths()["include"])
python_lib = Path(sysconfig.get_config_var("LIBDIR"))
# Match the deployment target that Python itself was built against
# unless the user overrides it. MLX's prebuilt wheels may require a
# higher minimum; in that case set MACOSX_DEPLOYMENT_TARGET explicitly.
deployment_target = os.environ.get(
"MACOSX_DEPLOYMENT_TARGET",
str(sysconfig.get_config_var("MACOSX_DEPLOYMENT_TARGET") or "11.0"),
)
def _python_eval(expr: str) -> str:
return subprocess.check_output(
[str(python_exe), "-c", expr], text=True
).strip()
nanobind_dir = Path(
_python_eval("import nanobind; print(nanobind.__path__[0])")
)
mlx_dir = Path(_python_eval("import mlx.core as mx; print(mx.__file__)"))
mlx_site = mlx_dir.parent
mlx_include = mlx_site / "include"
mlx_lib = mlx_site / "lib"
generated_dir = root / "build" / "metal"
generated_dir.mkdir(parents=True, exist_ok=True)
metallib_path = generated_dir / metallib_name
metal_std = os.environ.get("SGL_METAL_STD", "metal3.1")
ext_include_dirs = [Path(p) for p in (ext.include_dirs or [])]
host_includes = [
python_include,
nanobind_dir / "include",
nanobind_dir / "ext" / "robin_map" / "include",
mlx_include,
mlx_include / "metal_cpp",
]
all_includes = ext_include_dirs + host_includes
include_args = [f"-I{p}" for p in all_includes]
# `xcrun metal` accepts `-I` for header search; reuse the project
# include dirs so shaders can include shared MSL headers.
metal_include_args = [f"-I{p}" for p in ext_include_dirs]
if not metal_shader_sources:
raise RuntimeError("metal_shader_sources is empty; nothing to compile")
air_paths = []
for rel in metal_shader_sources:
metal_src = root / rel
if not metal_src.is_file():
raise RuntimeError(f"metal shader source not found: {metal_src}")
air_path = generated_dir / (metal_src.stem + ".air")
self.spawn(
[
"xcrun",
"-sdk",
"macosx",
"metal",
f"-std={metal_std}",
*metal_flags,
*metal_include_args,
"-c",
str(metal_src),
"-o",
str(air_path),
]
)
air_paths.append(str(air_path))
self.spawn(
[
"xcrun",
"-sdk",
"macosx",
"metallib",
*air_paths,
"-o",
str(metallib_path),
]
)
cflags = [
*cxx_flags,
f"-mmacosx-version-min={deployment_target}",
*include_args,
]
ldflags = [
"-shared",
"-undefined",
"dynamic_lookup",
f"-mmacosx-version-min={deployment_target}",
f"-L{python_lib}",
f"-L{mlx_lib}",
f"-Wl,-rpath,{mlx_lib}",
*[f"-l{lib}" for lib in libraries],
*[arg for fw in frameworks for arg in ("-framework", fw)],
]
objects = []
for src in ext.sources:
src_path = Path(src)
obj_path = generated_dir / (src_path.stem + ".o")
compile_cmd = [
*cxx_cmd,
*cflags,
"-c",
str(src_path),
"-o",
str(obj_path),
]
self.spawn(compile_cmd)
objects.append(str(obj_path))
nanobind_src = nanobind_dir / "src" / "nb_combined.cpp"
nanobind_obj = generated_dir / "nb_combined.o"
nanobind_cmd = [
*cxx_cmd,
*cflags,
"-DNB_COMPACT_ASSERTIONS",
"-DNB_BUILD",
"-DNB_SHARED",
"-c",
str(nanobind_src),
"-o",
str(nanobind_obj),
]
self.spawn(nanobind_cmd)
objects.append(str(nanobind_obj))
link_cmd = [
"c++",
*objects,
*ldflags,
"-o",
str(ext_path),
]
self.spawn(link_cmd)
# Stage the metallib next to the freshly-linked extension so that
# `install_lib` picks it up via `package_data={"sgl_kernel": ["*.metallib"]}`.
staged_metallib = ext_path.parent / metallib_path.name
if metallib_path.resolve() != staged_metallib.resolve():
shutil.copy2(metallib_path, staged_metallib)
ext_modules = [
Extension(
name=f"{operator_namespace}._metal",
sources=cxx_sources,
include_dirs=[str(p) for p in include_dirs],
language="c++",
)
]
setup(
name="sglang-kernel",
version=_get_version(),
packages=find_packages(where="python"),
package_dir={"": "python"},
package_data={"sgl_kernel": ["*.metallib"]},
include_package_data=True,
ext_modules=ext_modules,
cmdclass={"build_ext": BuildMetalExtension},
)