diff --git a/python/sglang/srt/arg_groups/choices.py b/python/sglang/srt/arg_groups/choices.py index 1ac5b914cb7e..828a5cf0ceb0 100644 --- a/python/sglang/srt/arg_groups/choices.py +++ b/python/sglang/srt/arg_groups/choices.py @@ -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", diff --git a/python/sglang/srt/arg_groups/platform_hook.py b/python/sglang/srt/arg_groups/platform_hook.py index 67007b73427f..7fdffb6c2aef 100644 --- a/python/sglang/srt/arg_groups/platform_hook.py +++ b/python/sglang/srt/arg_groups/platform_hook.py @@ -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. diff --git a/python/sglang/srt/arg_groups/validation_hook.py b/python/sglang/srt/arg_groups/validation_hook.py index 352b5f4c9624..5504e7bd9c58 100644 --- a/python/sglang/srt/arg_groups/validation_hook.py +++ b/python/sglang/srt/arg_groups/validation_hook.py @@ -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: @@ -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. diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index abf59fd1c0a6..b72d6a13de0f 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -31,6 +31,7 @@ is_hip, is_musa, is_npu, + is_xpu, ) if is_cuda(): @@ -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, @@ -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: @@ -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)