forked from Karylab-cklius/vllm
[XPU] use xpu topk topp sample kernel (#39285)
Signed-off-by: Kunshang Ji <kunshang.ji@intel.com>
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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"))
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user