From 2ceea429583ad6429a3bb12ba2e909c5cb99f059 Mon Sep 17 00:00:00 2001 From: Kunshang Ji Date: Tue, 5 May 2026 18:05:17 +0800 Subject: [PATCH] [XPU] use xpu topk topp sample kernel (#39285) Signed-off-by: Kunshang Ji --- vllm/_xpu_ops.py | 35 ++++++++++++++++++ vllm/envs.py | 5 +++ vllm/v1/sample/ops/topk_topp_sampler.py | 48 +++++++++++++++++++++++++ 3 files changed, 88 insertions(+) diff --git a/vllm/_xpu_ops.py b/vllm/_xpu_ops.py index 0b39a400012..09f700d0de7 100644 --- a/vllm/_xpu_ops.py +++ b/vllm/_xpu_ops.py @@ -185,6 +185,35 @@ def _xpu_ops_deepseek_scaling_rope_fake( return query, key +def _topk_topp_sample_impl( + random_sampled: torch.Tensor, + logits_to_return: torch.Tensor | None, + logits: torch.Tensor, + k: torch.Tensor | None, + p: torch.Tensor | None, + logprobs_mode: str, + seeds: torch.Tensor | None, + lambda_: float = 1.0, +) -> None: + torch.ops._xpu_C.topk_topp_sampler( + random_sampled, logits_to_return, logits, k, p, logprobs_mode, seeds, lambda_ + ) + return + + +def _topk_topp_sample_fake( + random_sampled: torch.Tensor, + logits_to_return: torch.Tensor | None, + logits: torch.Tensor, + k: torch.Tensor | None, + p: torch.Tensor | None, + logprobs_mode: str, + seeds: torch.Tensor | None, + lambda_: float = 1.0, +) -> None: + return + + def _xpu_mxfp8_quantize_impl( x: torch.Tensor, dtype: torch.dtype | None = None ) -> tuple[torch.Tensor, torch.Tensor]: @@ -691,6 +720,12 @@ class xpu_ops: fake_impl=_gdn_attention_core_xpu_fake, ) + direct_register_custom_op( + op_name="xpu_topk_topp_sampler", + op_func=_topk_topp_sample_impl, + fake_impl=_topk_topp_sample_fake, + ) + _OPS_REGISTERED = True diff --git a/vllm/envs.py b/vllm/envs.py index a5a4bfaffd5..ded474dc085 100755 --- a/vllm/envs.py +++ b/vllm/envs.py @@ -266,6 +266,7 @@ if TYPE_CHECKING: VLLM_MEMORY_PROFILER_ESTIMATE_CUDAGRAPHS: bool = True VLLM_NIXL_EP_MAX_NUM_RANKS: int = 32 VLLM_XPU_ENABLE_XPU_GRAPH: bool = False + VLLM_XPU_USE_SAMPLER_KERNEL: bool = True VLLM_LORA_ENABLE_DUAL_STREAM: bool = False @@ -1775,6 +1776,10 @@ environment_variables: dict[str, Callable[[], Any]] = { "VLLM_XPU_ENABLE_XPU_GRAPH": lambda: bool( int(os.getenv("VLLM_XPU_ENABLE_XPU_GRAPH", "0")) ), + # whether use xpu specific sample kernel + "VLLM_XPU_USE_SAMPLER_KERNEL": lambda: bool( + int(os.getenv("VLLM_XPU_USE_SAMPLER_KERNEL", "1")) + ), # Enable simple KV offload. "VLLM_USE_SIMPLE_KV_OFFLOAD": lambda: bool( int(os.getenv("VLLM_USE_SIMPLE_KV_OFFLOAD", "0")) diff --git a/vllm/v1/sample/ops/topk_topp_sampler.py b/vllm/v1/sample/ops/topk_topp_sampler.py index 70843be3969..363b113f0a4 100644 --- a/vllm/v1/sample/ops/topk_topp_sampler.py +++ b/vllm/v1/sample/ops/topk_topp_sampler.py @@ -82,6 +82,11 @@ class TopKTopPSampler(nn.Module): self.forward = self.forward_native else: self.forward = self.forward_cpu + elif current_platform.is_xpu(): + if envs.VLLM_XPU_USE_SAMPLER_KERNEL: + self.forward = self.forward_xpu + else: + self.forward = self.forward_native elif ( logprobs_mode not in ("processed_logits", "processed_logprobs") and rocm_aiter_ops.is_enabled() @@ -243,6 +248,49 @@ class TopKTopPSampler(nn.Module): return torch.multinomial(renorm_probs, num_samples=1).view(-1) raise RuntimeError("aiter_sample was called with no active top-k or top-p.") + def forward_xpu( + self, + logits: torch.Tensor, + generators: dict[int, torch.Generator], + k: torch.Tensor | None, + p: torch.Tensor | None, + ) -> tuple[torch.Tensor, torch.Tensor | None]: + if generators: + logger.warning_once( + "xpu kernel topk_topp_sampler does not support " + "per-request generators. Falling back to " + "PyTorch-native implementation." + ) + return self.forward_native(logits, generators, k, p) + random_sampled = torch.empty( + logits.shape[0], dtype=torch.int64, device=logits.device + ) + logits_to_return = None + if ( + self.logprobs_mode == "processed_logits" + or self.logprobs_mode == "processed_logprobs" + ): + logits_to_return = torch.empty_like(logits) + + assert len(generators) != logits.shape[0], ( + "xpu kernel topk_topp_sampler does not support batch-wise generators." + ) + generator = torch.xpu.default_generators[logits.device.index] + + state = generator.get_state() + seed, offset = state.view(torch.int64) + seeds = torch.tensor( + [seed, offset], dtype=torch.int64, device=torch.device("cpu") + ) + # The XPU kernel expects k as int64 (Long), but the input batch + # stores top_k as int32. Cast here to avoid dtype mismatch. + if k is not None: + k = k.to(torch.int64) + torch.ops.vllm.xpu_topk_topp_sampler( + random_sampled, logits_to_return, logits, k, p, self.logprobs_mode, seeds + ) + return random_sampled, logits_to_return + # Note: this is a workaround for # https://github.com/pytorch/pytorch/pull/151218