Skip to content

[XPU] Use sgl-kernel-xpu fused sampling kernels - #38510

Closed
KMS07 wants to merge 11 commits into
sgl-project:mainfrom
KMS07:xpu-sampling-backend
Closed

KMS07 wants to merge 11 commits into
sgl-project:mainfrom
KMS07:xpu-sampling-backend

Conversation

@KMS07

@KMS07 KMS07 commented Sep 8, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

On XPU, sampling_backend resolved to "pytorch" because _sampling_backend_default falls back whenever flashinfer is unavailable. The torch fallback sorts the full vocabulary per row on every decode step, so top-k/top-p sampling cost grows linearly with batch size.

sgl-kernel-xpu already ships the sampling ops, and they mirror the flashinfer API exactly (top_k_renorm_prob, top_p_renorm_prob, top_k_top_p_sampling_from_probs(..., filter_apply_order=...), min_p_sampling_from_probs), so the existing branch is reused as-is.

Adds an "intel_xpu" sampling backend, defaulted on --device xpu and guarded so an explicit --sampling-backend still wins.

Only the default, non-deterministic top-k/top-p/min-p path changes. Greedy and temperature-only requests are untouched, and --enable-deterministic-inference still resolves to "pytorch" since the XPU kernels take an RNG generator rather than the per-request sampling seed that batch-invariant sampling requires.

Modifications

  • Add "intel_xpu" to SAMPLING_BACKEND_CHOICES.
  • handle_xpu_backends resolves sampling_backend="intel_xpu" on --device xpu,
    guarded on is None so an explicit --sampling-backend still wins.
  • Import the four sampling ops from sgl_kernel under is_xpu(), alongside
    the existing is_musa() block.
  • Dispatch: if backend in ("flashinfer", "intel_xpu"). The signatures match
    (top_k_top_p_sampling_from_probs(probs, top_k, top_p, filter_apply_order=...),
    min_p_sampling_from_probs(probs, min_p)), so no separate branch is needed.
  • Reject --sampling-backend intel_xpu on a non-XPU device at startup
  • test/registered/sampling/test_sampling_mask.py runs on XPU. Device and
    fused-backend strings are derived once (get_device() / intel_xpu vs
    flashinfer) instead of hardcoded, so a single copy of each capture test
    covers CUDA, ROCm and XPU. This puts test_flashinfer_joint_cutoff_ties_match_capture
    on the XPU kernels: the fused SYCL joint sampler and the separately-written
    renorm kernels are verified to agree on cutoffs and ties, since a disagreement
    yields selected_weight == 0 and hence a -inf sampling logprob for a token
    that really was sampled. 10 capture tests pass on XPU.
  • Also fixes a latent portability bug this uncovered: TestSamplingMaskPacking
    hardcoded torch.cuda.Event(), which is a dummy class on XPU
    (RuntimeError: Tried to instantiate dummy base class Event). Now
    get_device_module().Event(), which resolves to torch.cuda.Event on
    CUDA/ROCm.

Accuracy Tests

End-to-end: GSM8K (Llama-3.3-70B-Instruct, TP=8, XPU)

Greedy (--temperature 0).

Backend Accuracy
pytorch 94.82
intel_xpu 95.20

Sampled (--temperature 0.7 --top-p 0.9, 1319 examples — the full GSM8K test split).
This is the configuration that actually exercises the new sampling kernels:

Backend Accuracy
pytorch 95.35
intel_xpu 95.28

min_p path: GSM8K (Llama-3.1-8B-Instruct, TP=4, XPU)

--temperature 0.7 --min-p 0.05 (top_p left at 1.0 to isolate min_p),
1319 examples — the full GSM8K test split. This is the three-kernel path
(top_k_renorm_prob -> top_p_renorm_prob -> min_p_sampling_from_probs)
rather than the fused joint call.

Backend Accuracy
pytorch 81.2
intel_xpu 81.2

Speed Tests and Profiling

Llama-3.3-70B-Instruct, TP=8 on Intel Arc Pro B60, --attention-backend intel_xpu,
sonnet dataset (1024 in / 1024 out, --ignore-eos), 320 prompts at concurrency 32,
--temperature 0.7 --top-p 0.9.

Both arms complete 320/320 requests with identical input (298,871) and output
(327,680) token counts, so the comparison is like-for-like.

Throughput

Metric pytorch intel_xpu Improvement
Output token throughput (tok/s) 368.95 413.49 +12.07%
Total token throughput (tok/s) 705.47 790.64 +12.07%
Request throughput (req/s) 0.3603 0.4038 +12.07%
Peak output tokens/s 480.0 544.0 +13.33%
Benchmark duration (s) 888.13 792.46 −10.77%

Per-token latency (TPOT)

Metric pytorch intel_xpu Improvement
Mean TPOT (ms) 83.49 74.27 −11.04%
Median TPOT (ms) 83.74 74.40 −11.15%
P90 TPOT (ms) 85.21 75.95 −10.87%
P99 TPOT (ms) 85.69 76.69 −10.50%

Inter-token latency (ITL)

Metric pytorch intel_xpu Improvement
Mean ITL (ms) 83.49 74.27 −11.04%
Median ITL (ms) 72.65 63.37 −12.77%
P90 ITL (ms) 73.84 64.57 −12.55%

End-to-end latency

Metric pytorch intel_xpu Improvement
Mean E2EL (ms) 88,757 79,194 −10.77%
Median E2EL (ms) 88,741 79,186 −10.77%
P90 E2EL (ms) 89,371 79,730 −10.79%
P99 E2EL (ms) 100,021 89,677 −10.34%

TTFT is essentially unchanged (median 2744.7 → 2717.6 ms, −0.99%), as expected:
sampling runs once per decode step and contributes nothing to prefill. The gain is
entirely in decode, which is where the median TPOT drop of 9.3 ms per step
shows up.

Reproduction

Server (baseline arm adds --sampling-backend pytorch; the xpu arm omits it and
resolves to "intel_xpu" automatically):

 python -m sglang.launch_server \
  --model meta-llama/Llama-3.3-70B-Instruct \
  --trust-remote-code --dtype bfloat16 --device xpu \
  --host 0.0.0.0 --port 8000 --tp-size 8 \
  --attention-backend intel_xpu \
  --context-length 9216 --chunked-prefill-size 2048 \
  --enable-mixed-chunk \
  --max-total-tokens 69120 --mem-fraction-static 0.90

Client — vLLM's bench serve uses the SGLang's OpenAI-compatible
endpoint. The sonnet dataset isnt available in the sglang bench serve.

vllm bench serve \
  --model meta-llama/Llama-3.3-70B-Instruct \
  --base-url http://localhost:8000 --backend vllm \
  --dataset-name sonnet \
  --dataset-path benchmarks/sonnet.txt \
  --sonnet-prefix-len 100 --sonnet-input-len 1024 --sonnet-output-len 1024 \
  --ignore-eos --trust-remote-code \
  --num-warmups 32 --num-prompts 320 --max-concurrency 32 \
  --percentile-metrics ttft,tpot,itl,e2el --metric-percentiles 90,99 \
  --seed 1 --save-result --temperature 0.7 --top-p 0.9

The min-p path uses three kernel launches (top_k_renorm_prob →
top_p_renorm_prob → min_p_sampling_from_probs) rather than the single
fused call the top-k/top-p path uses, so its speedup is lower — but it is
faster than the torch fallback at every batch size measured (1.93x at batch 1,
18.3x at batch 32, vocab 128256), so no special-casing is needed.

Checklist

Review and Merge Process

  1. Ping Merge Oncalls to start the process. See the PR Merge Process.
  2. Get approvals from CODEOWNERS and other reviewers.
  3. Trigger CI tests with comments or contact authorized users to do so.
    • Common commands include /tag-and-rerun-ci, /tag-run-ci-label, /rerun-failed-ci
  4. After green CI and required approvals, ask Merge Oncalls or people with Write permission to merge the PR.

CI States

Latest PR Test (Base): ❌ Run #35347617332
Latest PR Test (Extra): ❌ Run #35347617126
Latest PR Test (AMD ROCm 10): ❌ Run #35347617257

@KMS07

KMS07 commented Sep 8, 2026

Copy link
Copy Markdown
Contributor Author

Comment thread python/sglang/srt/arg_groups/choices.py Outdated
Comment thread python/sglang/srt/layers/sampler.py Outdated
@rbabukv

rbabukv commented Sep 9, 2026

Copy link
Copy Markdown

@siju-samuel pls review

Comment thread python/sglang/srt/layers/sampler.py Outdated
Comment thread python/sglang/srt/arg_groups/platform_hook.py Outdated
Comment thread test/registered/cpu/test_server_args_backend.py Outdated
Comment thread python/sglang/srt/layers/sampler.py Outdated
Comment thread python/sglang/srt/arg_groups/choices.py Outdated
@KMS07
KMS07 force-pushed the xpu-sampling-backend branch from cd481dc to 6aba4b3 Compare September 15, 2026 07:22
On XPU, `sampling_backend` resolved to "pytorch" because
`_sampling_backend_default` falls back whenever flashinfer is
unavailable. The torch fallback sorts the full vocabulary per row on
every decode step, so top-k/top-p sampling cost grows linearly with
batch size.

sgl-kernel-xpu already ships the sampling ops, and they mirror the
flashinfer API exactly (`top_k_renorm_prob`, `top_p_renorm_prob`,
`top_k_top_p_sampling_from_probs(..., filter_apply_order=...)`,
`min_p_sampling_from_probs`), so the existing branch is reused as-is.

Adds an "xpu" sampling backend, defaulted on `--device xpu` and guarded
so an explicit `--sampling-backend` still wins.

Only the default, non-deterministic top-k/top-p/min-p path changes.
Greedy and temperature-only requests are untouched, and
`--enable-deterministic-inference` still resolves to "pytorch" since
the XPU kernels take an RNG generator rather than the per-request
sampling seed that batch-invariant sampling requires.
@KMS07
KMS07 force-pushed the xpu-sampling-backend branch from cdf7a94 to 931c6cc Compare September 17, 2026 18:24
The fused-vs-renorm tie agreement was verified locally on XPU, but keep this
PR scoped to the sampling backend itself; the test wiring lands separately.
@KMS07
KMS07 marked this pull request as ready for review September 18, 2026 10:08
@polisettyvarma

Copy link
Copy Markdown
Contributor

/tag-run-ci-label

@github-actions github-actions Bot added the run-ci CI: run the baseline test suite on this PR label Sep 18, 2026
@KMS07
KMS07 marked this pull request as draft September 18, 2026 11:49
@KMS07
KMS07 marked this pull request as ready for review September 18, 2026 11:50
@KMS07

KMS07 commented Sep 18, 2026

Copy link
Copy Markdown
Contributor Author

@mingfeima pls review the changes

@KMS07

KMS07 commented Sep 28, 2026 •

Copy link
Copy Markdown
Contributor Author

PR 40664 merged the changes of this PR to main branch

@KMS07 KMS07 closed this Sep 28, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants