Skip to content
2 changes: 1 addition & 1 deletion python/sglang/srt/arg_groups/choices.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,7 +134,7 @@

GRAMMAR_BACKEND_CHOICES = ["xgrammar", "outlines", "llguidance", "none"]

SAMPLING_BACKEND_CHOICES = {"flashinfer", "pytorch", "ascend"}
SAMPLING_BACKEND_CHOICES = {"flashinfer", "pytorch", "ascend", "intel_xpu"}

MOE_RUNNER_BACKEND_CHOICES = [
"auto",
Expand Down
6 changes: 6 additions & 0 deletions python/sglang/srt/arg_groups/platform_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,6 +127,12 @@ def handle_symm_mem_device_support(server_args: Any):
def handle_xpu_backends(server_args: Any):
cfg = resolving_view(server_args)
if cfg.device == "xpu":
if cfg.sampling_backend is None:
declare_resolution(
server_args,
"_handle_xpu_backends",
sampling_backend="intel_xpu",
)
# Decode graph is opt-in on XPU: unless the user explicitly set
# --cuda-graph-backend-decode (or --cuda-graph-config), keep it
# disabled so the default startup doesn't require graph capture.
Expand Down
14 changes: 14 additions & 0 deletions python/sglang/srt/arg_groups/validation_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -266,6 +266,8 @@ def check_server_args(server_args: Any):
if cfg.enable_quant_communications and cfg.device != "npu":
raise ValueError("Communications quantization is only supported for NPU device")

validate_device_sampling_backend(cfg.sampling_backend, cfg.device)

# grpc_port is None for HTTP-only launches, so the == comparison is
# already False there; no explicit None check needed.
if not (cfg.smg_grpc_mode or cfg.grpc_mode) and cfg.grpc_port == cfg.port:
Expand Down Expand Up @@ -384,6 +386,18 @@ def check_load_publish_args(server_args: Any):
raise ValueError(reason)


def validate_device_sampling_backend(
sampling_backend: Optional[str], device: str
) -> None:
# sampler.py binds the intel_xpu kernels only under is_xpu(), so on another
# device the backend either aliases to flashinfer's names or NameErrors on
# the first non-greedy decode.
if sampling_backend == "intel_xpu" and device != "xpu":
raise ValueError(
f"--sampling-backend intel_xpu requires --device xpu, got --device {device}"
)


def validate_ib_devices(device_str: Optional[str]) -> Optional[str]:
"""
Validate IB devices before passing to mooncake.
Expand Down
9 changes: 5 additions & 4 deletions python/sglang/srt/layers/sampler.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
is_hip,
is_musa,
is_npu,
is_xpu,
)

if is_cuda():
Expand All @@ -43,7 +44,7 @@
top_p_renorm_prob,
)

if is_musa():
if is_musa() or is_xpu():
from sgl_kernel import (
min_p_sampling_from_probs,
top_k_renorm_prob,
Expand Down Expand Up @@ -73,7 +74,7 @@
SYNC_TOKEN_IDS_ACROSS_TP = get_bool_env_var("SYNC_TOKEN_IDS_ACROSS_TP")
SGLANG_RETURN_ORIGINAL_LOGPROB = get_bool_env_var("SGLANG_RETURN_ORIGINAL_LOGPROB")
_CUSTOM_SAMPLER_FACTORIES: Dict[str, Callable[[], "Sampler"]] = {}
_BUILT_IN_SAMPLING_BACKENDS = {"flashinfer", "pytorch", "ascend"}
_BUILT_IN_SAMPLING_BACKENDS = {"flashinfer", "pytorch", "ascend", "intel_xpu"}


def _trace_e2e_sampler(stage: str, **fields) -> None:
Expand Down Expand Up @@ -354,9 +355,9 @@ def _sample_from_probs(
)
else:
backend = get_exec().kernel.sampling_backend
if backend == "flashinfer":
if backend in ("flashinfer", "intel_xpu"):
assert sampling_info.sampling_seed is None, (
"Sampling seed is not supported for flashinfer backend"
f"Sampling seed is not supported for {backend} backend"
)
if sampling_info.need_min_p_sampling:
probs = top_k_renorm_prob(probs, sampling_info.top_ks)
Expand Down
Loading