[Apple Silicon] Add Metal kernel support in sgl-kernel (#23449)
Signed-off-by: Xiaodong Ye <yeahdongcn@gmail.com>
This commit is contained in:
@@ -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).
|
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
|
## Install SGLang
|
||||||
|
|
||||||
You can install SGLang using one of the methods below.
|
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
|
git clone https://github.com/sgl-project/sglang.git
|
||||||
cd sglang
|
cd sglang
|
||||||
|
|
||||||
# Install sglang python package
|
# Create and activate a virtual environment
|
||||||
pip install --upgrade pip
|
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
|
rm -f python/pyproject.toml && mv python/pyproject_other.toml python/pyproject.toml
|
||||||
uv pip install -e "python[all_mps]"
|
uv pip install -e "python[all_mps]"
|
||||||
```
|
```
|
||||||
|
|||||||
@@ -22,7 +22,7 @@ pip install --upgrade pip
|
|||||||
cd sgl-kernel
|
cd sgl-kernel
|
||||||
python setup_musa.py install
|
python setup_musa.py install
|
||||||
|
|
||||||
# Install sglang python package
|
# Install sglang python package along with diffusion support
|
||||||
cd ..
|
cd ..
|
||||||
rm -f python/pyproject.toml && mv python/pyproject_other.toml python/pyproject.toml
|
rm -f python/pyproject.toml && mv python/pyproject_other.toml python/pyproject.toml
|
||||||
pip install -e "python[all_musa]"
|
pip install -e "python[all_musa]"
|
||||||
|
|||||||
@@ -120,7 +120,7 @@ git clone https://github.com/sgl-project/sglang.git
|
|||||||
cd sglang
|
cd sglang
|
||||||
|
|
||||||
# Create and activate a virtual environment
|
# 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
|
source sglang-diffusion/bin/activate
|
||||||
|
|
||||||
# Install the Python packages
|
# Install the Python packages
|
||||||
|
|||||||
@@ -3,9 +3,10 @@
|
|||||||
# Install stubs early for platforms where certain dependencies are unavailable
|
# Install stubs early for platforms where certain dependencies are unavailable
|
||||||
# (e.g. macOS/MPS has no triton, and torch.mps lacks Stream / set_device /
|
# (e.g. macOS/MPS has no triton, and torch.mps lacks Stream / set_device /
|
||||||
# get_device_properties). This must run before any downstream imports.
|
# get_device_properties). This must run before any downstream imports.
|
||||||
|
import platform as _platform
|
||||||
import sys as _sys
|
import sys as _sys
|
||||||
|
|
||||||
if _sys.platform == "darwin":
|
if _sys.platform == "darwin" and _platform.machine() == "arm64":
|
||||||
try:
|
try:
|
||||||
import torch as _torch
|
import torch as _torch
|
||||||
|
|
||||||
@@ -22,6 +23,7 @@ if _sys.platform == "darwin":
|
|||||||
del _torch
|
del _torch
|
||||||
except ImportError:
|
except ImportError:
|
||||||
pass
|
pass
|
||||||
|
del _platform
|
||||||
del _sys
|
del _sys
|
||||||
|
|
||||||
from sglang.srt.utils.hf_transformers_patches import apply_all as _apply_hf_patches
|
from sglang.srt.utils.hf_transformers_patches import apply_all as _apply_hf_patches
|
||||||
|
|||||||
@@ -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.
|
||||||
@@ -1,212 +1,224 @@
|
|||||||
import torch
|
import platform
|
||||||
from sgl_kernel.debug_utils import maybe_wrap_debug_kernel
|
import sys
|
||||||
from sgl_kernel.load_utils import _load_architecture_specific_ops, _preload_cuda_library
|
|
||||||
|
|
||||||
# Initialize the ops library based on current GPU
|
from sgl_kernel.version import __version__ # noqa: F401
|
||||||
common_ops = _load_architecture_specific_ops()
|
|
||||||
|
|
||||||
# Preload the CUDA library to avoid the issue of libcudart.so.12 not found
|
# On macOS only the Metal extension is shipped; skip CUDA op loading and
|
||||||
if torch.version.cuda is not None:
|
# re-exports so those symbols are not exposed on Apple Silicon.
|
||||||
_preload_cuda_library()
|
if sys.platform == "darwin" and platform.machine() == "arm64":
|
||||||
|
from sgl_kernel.metal import *
|
||||||
|
else:
|
||||||
from sgl_kernel.allreduce import *
|
import torch
|
||||||
from sgl_kernel.attention import (
|
from sgl_kernel.debug_utils import maybe_wrap_debug_kernel
|
||||||
cutlass_mla_decode,
|
from sgl_kernel.load_utils import (
|
||||||
cutlass_mla_get_workspace_size,
|
_load_architecture_specific_ops,
|
||||||
merge_state_v2,
|
_preload_cuda_library,
|
||||||
)
|
|
||||||
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,
|
|
||||||
fused_add_rmsnorm,
|
|
||||||
gelu_and_mul,
|
|
||||||
gelu_tanh_and_mul,
|
|
||||||
gemma_fused_add_rmsnorm,
|
|
||||||
gemma_rmsnorm,
|
|
||||||
rmsnorm,
|
|
||||||
rotary_embedding,
|
|
||||||
silu_and_mul,
|
|
||||||
)
|
|
||||||
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 (
|
|
||||||
awq_dequantize,
|
|
||||||
bmm_fp8,
|
|
||||||
dsv3_fused_a_gemm,
|
|
||||||
dsv3_router_gemm,
|
|
||||||
fp8_blockwise_scaled_mm,
|
|
||||||
fp8_scaled_mm,
|
|
||||||
gptq_gemm,
|
|
||||||
gptq_shuffle,
|
|
||||||
int8_scaled_mm,
|
|
||||||
qserve_w4a8_per_chn_gemm,
|
|
||||||
qserve_w4a8_per_group_gemm,
|
|
||||||
sgl_per_token_group_quant_8bit,
|
|
||||||
sgl_per_token_group_quant_fp8,
|
|
||||||
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 (
|
|
||||||
transfer_kv_all_layer,
|
|
||||||
transfer_kv_all_layer_mla,
|
|
||||||
transfer_kv_per_layer,
|
|
||||||
transfer_kv_per_layer_mla,
|
|
||||||
)
|
|
||||||
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 (
|
|
||||||
apply_shuffle_mul_sum,
|
|
||||||
fp8_blockwise_scaled_grouped_mm,
|
|
||||||
fused_qk_norm_rope,
|
|
||||||
kimi_k2_moe_fused_gate,
|
|
||||||
moe_align_block_size,
|
|
||||||
moe_fused_gate,
|
|
||||||
moe_sum,
|
|
||||||
moe_sum_reduce,
|
|
||||||
prepare_moe_input,
|
|
||||||
topk_sigmoid,
|
|
||||||
topk_softmax,
|
|
||||||
)
|
|
||||||
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 (
|
|
||||||
top_k_renorm_prob,
|
|
||||||
top_p_renorm_prob,
|
|
||||||
)
|
|
||||||
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 (
|
|
||||||
fast_topk,
|
|
||||||
fast_topk_transform_fused,
|
|
||||||
fast_topk_transform_ragged_fused,
|
|
||||||
fast_topk_v2,
|
|
||||||
)
|
|
||||||
from sgl_kernel.version import __version__
|
|
||||||
|
|
||||||
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:
|
|
||||||
from sgl_kernel.musa import (
|
|
||||||
musa_batched_rotary_embedding_contiguous,
|
|
||||||
musa_fused_gemv,
|
|
||||||
musa_fused_moe_gemv,
|
|
||||||
musa_fused_mul_add,
|
|
||||||
musa_rotary_embedding_contiguous,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Initialize the ops library based on current GPU
|
||||||
|
common_ops = _load_architecture_specific_ops()
|
||||||
|
|
||||||
_DEBUG_EXPORT_NAMES = [
|
# Preload the CUDA library to avoid the issue of libcudart.so.12 not found
|
||||||
"apply_shuffle_mul_sum",
|
if torch.version.cuda is not None:
|
||||||
"apply_token_bitmask_inplace_cuda",
|
_preload_cuda_library()
|
||||||
"awq_dequantize",
|
|
||||||
"bmm_fp8",
|
|
||||||
"build_tree_kernel_efficient",
|
|
||||||
"causal_conv1d_fwd",
|
|
||||||
"causal_conv1d_update",
|
|
||||||
"concat_mla_absorb_q",
|
|
||||||
"concat_mla_k",
|
|
||||||
"copy_to_gpu_no_ce",
|
|
||||||
"cutlass_mla_decode",
|
|
||||||
"cutlass_mla_get_workspace_size",
|
|
||||||
"dsv3_fused_a_gemm",
|
|
||||||
"dsv3_router_gemm",
|
|
||||||
"es_fp8_blockwise_scaled_grouped_mm",
|
|
||||||
"es_sm100_mxfp8_blockscaled_grouped_mm",
|
|
||||||
"es_sm100_mxfp8_blockscaled_grouped_quant",
|
|
||||||
"fast_topk",
|
|
||||||
"fast_topk_transform_fused",
|
|
||||||
"fast_topk_transform_ragged_fused",
|
|
||||||
"fast_topk_v2",
|
|
||||||
"fp8_blockwise_scaled_grouped_mm",
|
|
||||||
"fp8_blockwise_scaled_mm",
|
|
||||||
"fp8_scaled_mm",
|
|
||||||
"fused_add_rmsnorm",
|
|
||||||
"fused_qk_norm_rope",
|
|
||||||
"gelu_and_mul",
|
|
||||||
"gelu_tanh_and_mul",
|
|
||||||
"gemma_fused_add_rmsnorm",
|
|
||||||
"gemma_rmsnorm",
|
|
||||||
"gptq_gemm",
|
|
||||||
"gptq_shuffle",
|
|
||||||
"int8_scaled_mm",
|
|
||||||
"kimi_k2_moe_fused_gate",
|
|
||||||
"merge_state_v2",
|
|
||||||
"moe_align_block_size",
|
|
||||||
"moe_fused_gate",
|
|
||||||
"moe_sum",
|
|
||||||
"moe_sum_reduce",
|
|
||||||
"prepare_moe_input",
|
|
||||||
"qserve_w4a8_per_chn_gemm",
|
|
||||||
"qserve_w4a8_per_group_gemm",
|
|
||||||
"reconstruct_indices_from_tree_mask",
|
|
||||||
"rmsnorm",
|
|
||||||
"rotary_embedding",
|
|
||||||
"segment_packbits",
|
|
||||||
"sgl_per_token_group_quant_8bit",
|
|
||||||
"sgl_per_token_group_quant_fp8",
|
|
||||||
"sgl_per_token_group_quant_int8",
|
|
||||||
"sgl_per_token_quant_fp8",
|
|
||||||
"shuffle_rows",
|
|
||||||
"silu_and_mul",
|
|
||||||
"top_k_renorm_prob",
|
|
||||||
"top_p_renorm_prob",
|
|
||||||
"topk_sigmoid",
|
|
||||||
"topk_softmax",
|
|
||||||
"transfer_kv_all_layer",
|
|
||||||
"transfer_kv_all_layer_mla",
|
|
||||||
"transfer_kv_per_layer",
|
|
||||||
"transfer_kv_per_layer_mla",
|
|
||||||
"tree_speculative_sampling_target_only",
|
|
||||||
"verify_tree_greedy",
|
|
||||||
"weak_ref_tensor",
|
|
||||||
]
|
|
||||||
|
|
||||||
if torch.version.hip is not None:
|
from sgl_kernel.allreduce import *
|
||||||
_DEBUG_EXPORT_NAMES.append("gelu_quick")
|
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 (
|
||||||
|
concat_mla_absorb_q,
|
||||||
|
concat_mla_k,
|
||||||
|
copy_to_gpu_no_ce,
|
||||||
|
fused_add_rmsnorm,
|
||||||
|
gelu_and_mul,
|
||||||
|
gelu_tanh_and_mul,
|
||||||
|
gemma_fused_add_rmsnorm,
|
||||||
|
gemma_rmsnorm,
|
||||||
|
rmsnorm,
|
||||||
|
rotary_embedding,
|
||||||
|
silu_and_mul,
|
||||||
|
)
|
||||||
|
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 (
|
||||||
|
awq_dequantize,
|
||||||
|
bmm_fp8,
|
||||||
|
dsv3_fused_a_gemm,
|
||||||
|
dsv3_router_gemm,
|
||||||
|
fp8_blockwise_scaled_mm,
|
||||||
|
fp8_scaled_mm,
|
||||||
|
gptq_gemm,
|
||||||
|
gptq_shuffle,
|
||||||
|
int8_scaled_mm,
|
||||||
|
qserve_w4a8_per_chn_gemm,
|
||||||
|
qserve_w4a8_per_group_gemm,
|
||||||
|
sgl_per_token_group_quant_8bit,
|
||||||
|
sgl_per_token_group_quant_fp8,
|
||||||
|
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 (
|
||||||
|
transfer_kv_all_layer,
|
||||||
|
transfer_kv_all_layer_mla,
|
||||||
|
transfer_kv_per_layer,
|
||||||
|
transfer_kv_per_layer_mla,
|
||||||
|
)
|
||||||
|
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 (
|
||||||
|
apply_shuffle_mul_sum,
|
||||||
|
fp8_blockwise_scaled_grouped_mm,
|
||||||
|
fused_qk_norm_rope,
|
||||||
|
kimi_k2_moe_fused_gate,
|
||||||
|
moe_align_block_size,
|
||||||
|
moe_fused_gate,
|
||||||
|
moe_sum,
|
||||||
|
moe_sum_reduce,
|
||||||
|
prepare_moe_input,
|
||||||
|
topk_sigmoid,
|
||||||
|
topk_softmax,
|
||||||
|
)
|
||||||
|
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 (
|
||||||
|
top_k_renorm_prob,
|
||||||
|
top_p_renorm_prob,
|
||||||
|
)
|
||||||
|
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 (
|
||||||
|
fast_topk,
|
||||||
|
fast_topk_transform_fused,
|
||||||
|
fast_topk_transform_ragged_fused,
|
||||||
|
fast_topk_v2,
|
||||||
|
)
|
||||||
|
from sgl_kernel.version import __version__
|
||||||
|
|
||||||
for _name in _DEBUG_EXPORT_NAMES:
|
if torch.version.hip is not None:
|
||||||
if _name in globals():
|
from sgl_kernel.elementwise import gelu_quick
|
||||||
globals()[_name] = maybe_wrap_debug_kernel(
|
|
||||||
globals()[_name], f"sgl_kernel.{_name}"
|
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,
|
||||||
|
musa_fused_moe_gemv,
|
||||||
|
musa_fused_mul_add,
|
||||||
|
musa_rotary_embedding_contiguous,
|
||||||
)
|
)
|
||||||
|
|
||||||
del _name
|
_DEBUG_EXPORT_NAMES = [
|
||||||
del _DEBUG_EXPORT_NAMES
|
"apply_shuffle_mul_sum",
|
||||||
|
"apply_token_bitmask_inplace_cuda",
|
||||||
|
"awq_dequantize",
|
||||||
|
"bmm_fp8",
|
||||||
|
"build_tree_kernel_efficient",
|
||||||
|
"causal_conv1d_fwd",
|
||||||
|
"causal_conv1d_update",
|
||||||
|
"concat_mla_absorb_q",
|
||||||
|
"concat_mla_k",
|
||||||
|
"copy_to_gpu_no_ce",
|
||||||
|
"cutlass_mla_decode",
|
||||||
|
"cutlass_mla_get_workspace_size",
|
||||||
|
"dsv3_fused_a_gemm",
|
||||||
|
"dsv3_router_gemm",
|
||||||
|
"es_fp8_blockwise_scaled_grouped_mm",
|
||||||
|
"es_sm100_mxfp8_blockscaled_grouped_mm",
|
||||||
|
"es_sm100_mxfp8_blockscaled_grouped_quant",
|
||||||
|
"fast_topk",
|
||||||
|
"fast_topk_transform_fused",
|
||||||
|
"fast_topk_transform_ragged_fused",
|
||||||
|
"fast_topk_v2",
|
||||||
|
"fp8_blockwise_scaled_grouped_mm",
|
||||||
|
"fp8_blockwise_scaled_mm",
|
||||||
|
"fp8_scaled_mm",
|
||||||
|
"fused_add_rmsnorm",
|
||||||
|
"fused_qk_norm_rope",
|
||||||
|
"gelu_and_mul",
|
||||||
|
"gelu_tanh_and_mul",
|
||||||
|
"gemma_fused_add_rmsnorm",
|
||||||
|
"gemma_rmsnorm",
|
||||||
|
"gptq_gemm",
|
||||||
|
"gptq_shuffle",
|
||||||
|
"int8_scaled_mm",
|
||||||
|
"kimi_k2_moe_fused_gate",
|
||||||
|
"merge_state_v2",
|
||||||
|
"moe_align_block_size",
|
||||||
|
"moe_fused_gate",
|
||||||
|
"moe_sum",
|
||||||
|
"moe_sum_reduce",
|
||||||
|
"prepare_moe_input",
|
||||||
|
"qserve_w4a8_per_chn_gemm",
|
||||||
|
"qserve_w4a8_per_group_gemm",
|
||||||
|
"reconstruct_indices_from_tree_mask",
|
||||||
|
"rmsnorm",
|
||||||
|
"rotary_embedding",
|
||||||
|
"segment_packbits",
|
||||||
|
"sgl_per_token_group_quant_8bit",
|
||||||
|
"sgl_per_token_group_quant_fp8",
|
||||||
|
"sgl_per_token_group_quant_int8",
|
||||||
|
"sgl_per_token_quant_fp8",
|
||||||
|
"shuffle_rows",
|
||||||
|
"silu_and_mul",
|
||||||
|
"top_k_renorm_prob",
|
||||||
|
"top_p_renorm_prob",
|
||||||
|
"topk_sigmoid",
|
||||||
|
"topk_softmax",
|
||||||
|
"transfer_kv_all_layer",
|
||||||
|
"transfer_kv_all_layer_mla",
|
||||||
|
"transfer_kv_per_layer",
|
||||||
|
"transfer_kv_per_layer_mla",
|
||||||
|
"tree_speculative_sampling_target_only",
|
||||||
|
"verify_tree_greedy",
|
||||||
|
"weak_ref_tensor",
|
||||||
|
]
|
||||||
|
|
||||||
|
if torch.version.hip is not None:
|
||||||
|
_DEBUG_EXPORT_NAMES.append("gelu_quick")
|
||||||
|
|
||||||
def create_greenctx_stream_by_value(*args, **kwargs):
|
for _name in _DEBUG_EXPORT_NAMES:
|
||||||
from sgl_kernel.spatial import create_greenctx_stream_by_value as _impl
|
if _name in globals():
|
||||||
|
globals()[_name] = maybe_wrap_debug_kernel(
|
||||||
|
globals()[_name], f"sgl_kernel.{_name}"
|
||||||
|
)
|
||||||
|
|
||||||
return _impl(*args, **kwargs)
|
del _name
|
||||||
|
del _DEBUG_EXPORT_NAMES
|
||||||
|
|
||||||
|
def create_greenctx_stream_by_value(*args, **kwargs):
|
||||||
|
from sgl_kernel.spatial import create_greenctx_stream_by_value as _impl
|
||||||
|
|
||||||
def get_sm_available(*args, **kwargs):
|
return _impl(*args, **kwargs)
|
||||||
from sgl_kernel.spatial import get_sm_available as _impl
|
|
||||||
|
|
||||||
return _impl(*args, **kwargs)
|
def get_sm_available(*args, **kwargs):
|
||||||
|
from sgl_kernel.spatial import get_sm_available as _impl
|
||||||
|
|
||||||
|
return _impl(*args, **kwargs)
|
||||||
|
|||||||
@@ -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.
|
||||||
@@ -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},
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user