Repository navigation
Conversation
Contributor
Author
|
@siju-samuel pls review |
siju-samuel
suggested changes
Sep 9, 2026
KMS07
force-pushed
the
xpu-sampling-backend
branch
from
September 15, 2026 07:22
cd481dc to
6aba4b3
Compare
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
force-pushed
the
xpu-sampling-backend
branch
from
September 17, 2026 18:24
cdf7a94 to
931c6cc
Compare
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.
rahulvijayaraghavan
approved these changes
Sep 18, 2026
siju-samuel
approved these changes
Sep 18, 2026
KMS07
marked this pull request as ready for review
September 18, 2026 10:08
KMS07
requested review from
BBuf,
Edwardf0t1,
Fridge003,
HaiShaw,
Ying1123,
ch-wan,
ispobock and
merrymercy
as code owners
September 18, 2026 10:08
Contributor
|
/tag-run-ci-label |
KMS07
marked this pull request as draft
September 18, 2026 11:49
KMS07
marked this pull request as ready for review
September 18, 2026 11:50
Contributor
Author
|
@mingfeima pls review the changes |
5 tasks
airMeng
pushed a commit
to airMeng/sglang
that referenced
this pull request
Sep 22, 2026
airMeng
pushed a commit
to airMeng/sglang
that referenced
this pull request
Sep 24, 2026
2 tasks done
Contributor
Author
|
PR 40664 merged the changes of this PR to |
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Motivation
On XPU,
sampling_backendresolved to "pytorch" because_sampling_backend_defaultfalls 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 xpuand guarded so an explicit--sampling-backendstill wins.Only the default, non-deterministic top-k/top-p/min-p path changes. Greedy and temperature-only requests are untouched, and
--enable-deterministic-inferencestill resolves to "pytorch" since the XPU kernels take an RNG generator rather than the per-request sampling seed that batch-invariant sampling requires.Modifications
"intel_xpu"toSAMPLING_BACKEND_CHOICES.handle_xpu_backendsresolvessampling_backend="intel_xpu"on--device xpu,guarded on
is Noneso an explicit--sampling-backendstill wins.sgl_kernelunderis_xpu(), alongsidethe existing
is_musa()block.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.--sampling-backend intel_xpuon a non-XPU device at startuptest/registered/sampling/test_sampling_mask.pyruns on XPU. Device andfused-backend strings are derived once (
get_device()/intel_xpuvsflashinfer) instead of hardcoded, so a single copy of each capture testcovers CUDA, ROCm and XPU. This puts
test_flashinfer_joint_cutoff_ties_match_captureon 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 == 0and hence a-infsampling logprob for a tokenthat really was sampled. 10 capture tests pass on XPU.
TestSamplingMaskPackinghardcoded
torch.cuda.Event(), which is a dummy class on XPU(
RuntimeError: Tried to instantiate dummy base class Event). Nowget_device_module().Event(), which resolves totorch.cuda.EventonCUDA/ROCm.
Accuracy Tests
End-to-end: GSM8K (Llama-3.3-70B-Instruct, TP=8, XPU)
Greedy (
--temperature 0).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:
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.
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
Per-token latency (TPOT)
Inter-token latency (ITL)
End-to-end latency
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 andresolves to
"intel_xpu"automatically):Client — vLLM's
bench serveuses the SGLang's OpenAI-compatibleendpoint. The sonnet dataset isnt available in the sglang bench serve.
The min-p path uses three kernel launches (
top_k_renorm_prob→top_p_renorm_prob→min_p_sampling_from_probs) rather than the singlefused 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
/tag-and-rerun-ci,/tag-run-ci-label,/rerun-failed-ciCI States
Latest PR Test (Base): ❌ Run #35347617332
Latest PR Test (Extra): ❌ Run #35347617126
Latest PR Test (AMD ROCm 10): ❌ Run #35347617257