[XPU] use xpu topk topp sample kernel (#39285)

Signed-off-by: Kunshang Ji <kunshang.ji@intel.com>
This commit is contained in:
Kunshang Ji
2026-05-05 18:05:17 +08:00
committed by GitHub
parent bee126165f
commit 2ceea42958
3 changed files with 88 additions and 0 deletions
+35
View File
@@ -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
+5
View File
@@ -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"))
+48
View File
@@ -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