Speculative decoding support on XPU (#23180)
This commit is contained in:
@@ -0,0 +1,106 @@
|
||||
# Copyright 2023-2026 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""XPUCudaGraphBackend — Intel XPU full-graph capture (torch.xpu.XPUGraph)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import contextmanager
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.model_executor.runner_backend.base_cuda_graph_backend import (
|
||||
BaseCudaGraphBackend,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.model_executor.runner.base_cuda_graph_runner import (
|
||||
BaseCudaGraphRunner,
|
||||
)
|
||||
|
||||
|
||||
class XPUCudaGraphBackend(BaseCudaGraphBackend):
|
||||
"""One torch.xpu.XPUGraph per shape; attention metadata is
|
||||
captured inside the graph.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
cuda_graph_runner: BaseCudaGraphRunner,
|
||||
) -> None:
|
||||
self._graphs: Dict[Any, torch.xpu.XPUGraph] = {}
|
||||
self._outputs: Dict[Any, Any] = {}
|
||||
self._pool = None
|
||||
self._device_module = cuda_graph_runner.device_module
|
||||
self._tp_group = cuda_graph_runner.model_runner.tp_group
|
||||
self._capture_stream: Optional[torch.xpu.Stream] = None
|
||||
|
||||
@contextmanager
|
||||
def capture_session(self, stream: torch.xpu.Stream):
|
||||
if self._pool is None:
|
||||
self._pool = self._device_module.graph_pool_handle()
|
||||
self._capture_stream = stream
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
self._capture_stream = None
|
||||
|
||||
def capture_one(
|
||||
self,
|
||||
shape_key: Any,
|
||||
forward_fn: Callable[[], Any],
|
||||
dummies: Optional[Any] = None,
|
||||
post_warmup_hook: Optional[Callable[[], None]] = None,
|
||||
) -> None:
|
||||
# Two warmups so kernels are loaded and one-time setup is paid before capture.
|
||||
# post_warmup_hook lets the attention backend reset state that warmup mutated.
|
||||
for _ in range(2):
|
||||
self._device_module.synchronize()
|
||||
self._tp_group.barrier()
|
||||
forward_fn()
|
||||
if post_warmup_hook is not None:
|
||||
post_warmup_hook()
|
||||
|
||||
graph = torch.xpu.XPUGraph()
|
||||
|
||||
# graph_ctx: Callable[..., AbstractContextManager]
|
||||
graph_ctx = self._device_module.graph
|
||||
|
||||
with graph_ctx(graph, pool=self._pool, stream=self._capture_stream):
|
||||
out = forward_fn()
|
||||
|
||||
self._graphs[shape_key] = graph
|
||||
self._outputs[shape_key] = out
|
||||
|
||||
def can_run(self, forward_batch: ForwardBatch, shape_key: Any) -> bool:
|
||||
return shape_key in self._graphs
|
||||
|
||||
@contextmanager
|
||||
def replay_session(self):
|
||||
yield
|
||||
|
||||
def replay(
|
||||
self,
|
||||
shape_key: Any,
|
||||
static_forward_batch: ForwardBatch,
|
||||
**kwargs,
|
||||
) -> Any:
|
||||
self._graphs[shape_key].replay()
|
||||
return self._outputs[shape_key]
|
||||
|
||||
def cleanup(self) -> None:
|
||||
self._graphs.clear()
|
||||
self._outputs.clear()
|
||||
self._pool = None
|
||||
@@ -71,6 +71,12 @@ def resolve_decode_backend(
|
||||
return NPUCudaGraphBackend(
|
||||
cuda_graph_runner, enable_memory_saver=enable_memory_saver
|
||||
)
|
||||
elif model_runner.device == "xpu":
|
||||
from sglang.srt.hardware_backend.xpu.xpu_cudagraph_backend import (
|
||||
XPUCudaGraphBackend,
|
||||
)
|
||||
|
||||
return XPUCudaGraphBackend(cuda_graph_runner)
|
||||
|
||||
if backend_name == Backend.BREAKABLE:
|
||||
return BreakableCudaGraphBackend(
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import math
|
||||
from collections import defaultdict
|
||||
from enum import IntEnum
|
||||
@@ -19,7 +20,17 @@ from sglang.srt.mem_cache.common import (
|
||||
get_alloc_reserve_per_decode,
|
||||
get_last_loc,
|
||||
)
|
||||
from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu
|
||||
from sglang.srt.speculative.triton_ops.spec_tree import (
|
||||
sgl_build_tree_kernel_efficient_triton,
|
||||
verify_tree_greedy_kernel_triton,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
is_cuda,
|
||||
is_hip,
|
||||
is_musa,
|
||||
is_npu,
|
||||
is_xpu,
|
||||
)
|
||||
from sglang.srt.utils.async_probe import maybe_detect_oob
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -34,6 +45,9 @@ _is_cuda = is_cuda()
|
||||
_is_hip = is_hip()
|
||||
_is_npu = is_npu()
|
||||
_is_musa = is_musa()
|
||||
_is_xpu = is_xpu()
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
if _is_cuda or _is_hip or _is_musa:
|
||||
from sgl_kernel import (
|
||||
@@ -214,6 +228,21 @@ def build_tree_kernel_efficient(
|
||||
num_verify_tokens,
|
||||
tree_mask_mode,
|
||||
)
|
||||
elif _is_xpu:
|
||||
sgl_build_tree_kernel_triton(
|
||||
parent_list,
|
||||
top_scores_index,
|
||||
seq_lens,
|
||||
tree_mask,
|
||||
positions,
|
||||
retrieve_index,
|
||||
retrieve_next_token,
|
||||
retrieve_next_sibling,
|
||||
topk,
|
||||
spec_steps,
|
||||
num_verify_tokens,
|
||||
tree_mask_mode,
|
||||
)
|
||||
else:
|
||||
sgl_build_tree_kernel_efficient(
|
||||
parent_list,
|
||||
@@ -239,6 +268,88 @@ def build_tree_kernel_efficient(
|
||||
)
|
||||
|
||||
|
||||
def sgl_build_tree_kernel_triton(
|
||||
parent_list: torch.Tensor,
|
||||
selected_index: torch.Tensor,
|
||||
verified_seq_len: torch.Tensor,
|
||||
tree_mask: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
retrieve_index: torch.Tensor,
|
||||
retrieve_next_token: torch.Tensor,
|
||||
retrieve_next_sibling: torch.Tensor,
|
||||
topk: int,
|
||||
depth: int,
|
||||
draft_token_num: int,
|
||||
tree_mask_mode: TreeMaskMode = TreeMaskMode.FULL_MASK,
|
||||
):
|
||||
"""Triton-based implementation."""
|
||||
# TODO: Add support for QLEN_ONLY_BITPACKING mode
|
||||
if tree_mask_mode == TreeMaskMode.QLEN_ONLY_BITPACKING:
|
||||
raise NotImplementedError(
|
||||
"QLEN_ONLY_BITPACKING is not supported in Triton implementation"
|
||||
)
|
||||
|
||||
batch_size = verified_seq_len.shape[0]
|
||||
seq_len_prefix_sum = torch.cumsum(verified_seq_len, dim=0) - verified_seq_len
|
||||
|
||||
# Launch kernel with one program per batch item
|
||||
grid = (batch_size,)
|
||||
|
||||
sgl_build_tree_kernel_efficient_triton[grid](
|
||||
parent_list,
|
||||
selected_index,
|
||||
verified_seq_len,
|
||||
seq_len_prefix_sum,
|
||||
tree_mask,
|
||||
positions,
|
||||
retrieve_index,
|
||||
retrieve_next_token,
|
||||
retrieve_next_sibling,
|
||||
topk=topk,
|
||||
depth=depth,
|
||||
draft_token_num=draft_token_num,
|
||||
tree_mask_mode=int(tree_mask_mode),
|
||||
batch_size=batch_size,
|
||||
parent_list_stride=(
|
||||
parent_list.stride(0) if parent_list.dim() > 1 else parent_list.shape[0]
|
||||
),
|
||||
selected_index_stride=selected_index.stride(0),
|
||||
)
|
||||
|
||||
|
||||
def verify_tree_greedy_triton(
|
||||
predicts: torch.Tensor,
|
||||
accept_index: torch.Tensor,
|
||||
accept_token_num: torch.Tensor,
|
||||
candidates: torch.Tensor,
|
||||
retrieve_index: torch.Tensor,
|
||||
retrieve_next_token: torch.Tensor,
|
||||
retrieve_next_sibling: torch.Tensor,
|
||||
target_predict: torch.Tensor,
|
||||
):
|
||||
"""Triton-based implementation."""
|
||||
batch_size = candidates.shape[0]
|
||||
num_speculative_tokens = accept_index.shape[1]
|
||||
num_draft_tokens = candidates.shape[1]
|
||||
|
||||
# Launch kernel with one program per batch item
|
||||
grid = (batch_size,)
|
||||
|
||||
verify_tree_greedy_kernel_triton[grid](
|
||||
predicts,
|
||||
accept_index,
|
||||
accept_token_num,
|
||||
candidates,
|
||||
retrieve_index,
|
||||
retrieve_next_token,
|
||||
retrieve_next_sibling,
|
||||
target_predict,
|
||||
batch_size=batch_size,
|
||||
num_speculative_tokens=num_speculative_tokens,
|
||||
num_draft_tokens=num_draft_tokens,
|
||||
)
|
||||
|
||||
|
||||
def verify_tree_greedy_func(
|
||||
predicts: torch.Tensor,
|
||||
accept_index: torch.Tensor,
|
||||
@@ -279,6 +390,17 @@ def verify_tree_greedy_func(
|
||||
retrive_next_sibling=retrieve_next_sibling,
|
||||
target_predict=target_predict,
|
||||
)
|
||||
elif _is_xpu:
|
||||
verify_tree_greedy_triton(
|
||||
predicts=predicts,
|
||||
accept_index=accept_index,
|
||||
accept_token_num=accept_token_num,
|
||||
candidates=candidates,
|
||||
retrieve_index=retrieve_index,
|
||||
retrieve_next_token=retrieve_next_token,
|
||||
retrieve_next_sibling=retrieve_next_sibling,
|
||||
target_predict=target_predict,
|
||||
)
|
||||
return predicts, accept_index, accept_token_num
|
||||
|
||||
|
||||
@@ -495,7 +617,7 @@ def eagle_sample(
|
||||
|
||||
# Sample tokens
|
||||
target_predict = None
|
||||
if sampling_info.is_all_greedy or _is_npu or _is_hip:
|
||||
if sampling_info.is_all_greedy or _is_npu or _is_hip or _is_xpu:
|
||||
target_predict = torch.argmax(next_token_logits, dim=-1)
|
||||
target_predict = target_predict.reshape(bs, verify_input.draft_token_num)
|
||||
predict, accept_index, num_correct_drafts = verify_tree_greedy_func(
|
||||
|
||||
@@ -103,6 +103,7 @@ from sglang.srt.utils.common import (
|
||||
is_hip,
|
||||
is_musa,
|
||||
is_npu,
|
||||
is_xpu,
|
||||
log_info_on_rank0,
|
||||
)
|
||||
from sglang.srt.utils.patch_torch import monkey_patch_torch_reductions
|
||||
@@ -111,6 +112,7 @@ _is_npu = is_npu()
|
||||
_is_cuda = is_cuda()
|
||||
_is_musa = is_musa()
|
||||
_is_hip = is_hip()
|
||||
_is_xpu = is_xpu()
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -377,6 +379,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
||||
return
|
||||
|
||||
Device2DraftCudaGraphRunner = {
|
||||
"xpu": EAGLEDraftCudaGraphRunner,
|
||||
"npu": EAGLEDraftNpuGraphRunner,
|
||||
"cuda": EAGLEDraftCudaGraphRunner,
|
||||
"musa": EAGLEDraftCudaGraphRunner,
|
||||
@@ -406,6 +409,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
||||
)
|
||||
|
||||
Device2ExtendCudaGraphRunner = {
|
||||
"xpu": EAGLEDraftExtendCudaGraphRunner,
|
||||
"npu": EAGLEDraftExtendNpuGraphRunner,
|
||||
"cuda": EAGLEDraftExtendCudaGraphRunner,
|
||||
"musa": EAGLEDraftCudaGraphRunner,
|
||||
@@ -448,6 +452,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
||||
# TODO: support draft extend cuda graph for more attention backends
|
||||
if self.draft_extend_attn_backend and (
|
||||
_is_npu
|
||||
or _is_xpu
|
||||
or supports_cuda_draft_extend_graph
|
||||
or supports_hip_aiter_draft_extend_graph
|
||||
):
|
||||
|
||||
@@ -43,7 +43,7 @@ from sglang.srt.speculative.triton_ops.cache_locs import (
|
||||
from sglang.srt.speculative.triton_ops.eagle import (
|
||||
fill_accept_out_cache_loc as fill_accept_out_cache_loc,
|
||||
)
|
||||
from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu, next_power_of_2
|
||||
from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu, is_xpu, next_power_of_2
|
||||
from sglang.srt.utils.async_probe import maybe_detect_oob
|
||||
from sglang.srt.utils.nvtx_utils import profile_range
|
||||
|
||||
@@ -51,6 +51,7 @@ _is_cuda = is_cuda()
|
||||
_is_hip = is_hip()
|
||||
_is_npu = is_npu()
|
||||
_is_musa = is_musa()
|
||||
_is_xpu = is_xpu()
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject
|
||||
@@ -189,7 +190,7 @@ def spec_need_hidden_states(server_args: Optional[ServerArgs] = None) -> bool:
|
||||
return not server_args.enable_multi_layer_eagle
|
||||
|
||||
|
||||
@torch.compile(dynamic=True, disable=_is_npu)
|
||||
@torch.compile(dynamic=True, disable=_is_npu or _is_xpu)
|
||||
def create_num_accept_tokens_filter(
|
||||
num_correct_drafts: torch.Tensor,
|
||||
unfinished_index_device: torch.Tensor,
|
||||
@@ -223,7 +224,7 @@ def _select_top_k_tokens_first(
|
||||
return input_ids, hidden_states, topk_p, tree_info
|
||||
|
||||
|
||||
@torch.compile(dynamic=True, disable=_is_npu)
|
||||
@torch.compile(dynamic=True, disable=_is_npu or _is_xpu)
|
||||
def _select_top_k_tokens_later(
|
||||
i: int,
|
||||
topk_p: torch.Tensor,
|
||||
|
||||
@@ -4,12 +4,13 @@ import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu, next_power_of_2
|
||||
from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu, is_xpu, next_power_of_2
|
||||
|
||||
_is_cuda = is_cuda()
|
||||
_is_hip = is_hip()
|
||||
_is_npu = is_npu()
|
||||
_is_musa = is_musa()
|
||||
_is_xpu = is_xpu()
|
||||
|
||||
|
||||
@triton.jit
|
||||
@@ -343,7 +344,7 @@ def assign_extend_cache_locs_func(
|
||||
draft_token_num: int,
|
||||
device,
|
||||
) -> torch.Tensor:
|
||||
if _is_cuda or _is_hip or _is_musa:
|
||||
if _is_cuda or _is_hip or _is_musa or _is_xpu:
|
||||
out_cache_loc = torch.empty(
|
||||
(batch_size * draft_token_num,),
|
||||
dtype=torch.int64,
|
||||
|
||||
@@ -0,0 +1,281 @@
|
||||
# Copyright 2023-2026 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@triton.jit
|
||||
def sgl_build_tree_kernel_efficient_triton(
|
||||
parent_list_ptr,
|
||||
selected_index_ptr,
|
||||
verified_seq_len_ptr,
|
||||
seq_len_prefix_sum_ptr,
|
||||
tree_mask_ptr,
|
||||
positions_ptr,
|
||||
retrieve_index_ptr,
|
||||
retrieve_next_token_ptr,
|
||||
retrieve_next_sibling_ptr,
|
||||
topk: tl.constexpr,
|
||||
depth: tl.constexpr,
|
||||
draft_token_num: tl.constexpr,
|
||||
tree_mask_mode: tl.constexpr,
|
||||
batch_size: tl.constexpr,
|
||||
parent_list_stride: tl.constexpr,
|
||||
selected_index_stride: tl.constexpr,
|
||||
):
|
||||
"""
|
||||
Triton kernel for building EAGLE tree structure.
|
||||
Each program handles one batch item (batch_idx).
|
||||
"""
|
||||
batch_idx = tl.program_id(0)
|
||||
|
||||
# Calculate seq_tree_idx
|
||||
seq_len = tl.load(verified_seq_len_ptr + batch_idx)
|
||||
seq_len_prefix_sum = tl.load(seq_len_prefix_sum_ptr + batch_idx)
|
||||
|
||||
# Cast initial value to match the dtype of loaded tensors to avoid type inconsistency
|
||||
seq_tree_idx = (
|
||||
tl.cast(draft_token_num * draft_token_num * batch_idx, seq_len.dtype)
|
||||
+ seq_len_prefix_sum * draft_token_num
|
||||
)
|
||||
|
||||
positions_offset = batch_idx * draft_token_num
|
||||
tl.store(positions_ptr + positions_offset, seq_len)
|
||||
|
||||
retrieve_index_offset = batch_idx * draft_token_num
|
||||
|
||||
# Build retrieval index structure (reverse loop from draft_token_num-1 to 1)
|
||||
for i in range(draft_token_num - 1, 0, -1):
|
||||
current_token_idx = retrieve_index_offset + i
|
||||
tl.store(
|
||||
retrieve_index_ptr + batch_idx * draft_token_num + i,
|
||||
current_token_idx,
|
||||
)
|
||||
|
||||
parent_tb_idx = (
|
||||
tl.load(selected_index_ptr + batch_idx * selected_index_stride + (i - 1))
|
||||
// topk
|
||||
)
|
||||
parent_position = 0
|
||||
found = 0
|
||||
|
||||
if parent_tb_idx == 0:
|
||||
found = 1
|
||||
else:
|
||||
parent_token_idx = tl.load(
|
||||
parent_list_ptr + batch_idx * parent_list_stride + parent_tb_idx
|
||||
)
|
||||
|
||||
# Find parent position
|
||||
for pp in range(draft_token_num - 1):
|
||||
if found == 0:
|
||||
sel_idx = tl.load(
|
||||
selected_index_ptr + batch_idx * selected_index_stride + pp
|
||||
)
|
||||
if sel_idx == parent_token_idx:
|
||||
parent_position = pp + 1
|
||||
found = 1
|
||||
|
||||
if found == 1:
|
||||
# Update next token links
|
||||
next_tok_addr = (
|
||||
retrieve_next_token_ptr + batch_idx * draft_token_num + parent_position
|
||||
)
|
||||
next_tok = tl.load(next_tok_addr)
|
||||
|
||||
if next_tok == -1:
|
||||
tl.store(next_tok_addr, i)
|
||||
else:
|
||||
tl.store(next_tok_addr, i)
|
||||
tl.store(
|
||||
retrieve_next_sibling_ptr + batch_idx * draft_token_num + i,
|
||||
next_tok,
|
||||
)
|
||||
|
||||
tl.store(retrieve_index_ptr + batch_idx * draft_token_num, retrieve_index_offset)
|
||||
|
||||
# Process all draft token indices for tree mask
|
||||
for draft_tokenx in range(draft_token_num):
|
||||
if tree_mask_mode == 0: # FULL_MASK
|
||||
token_tree_idx = (
|
||||
seq_tree_idx + (seq_len + draft_token_num) * draft_tokenx + seq_len + 1
|
||||
)
|
||||
else:
|
||||
token_tree_idx = (
|
||||
draft_token_num * draft_token_num * batch_idx
|
||||
+ draft_token_num * draft_tokenx
|
||||
+ 1
|
||||
)
|
||||
|
||||
tl.store(tree_mask_ptr + token_tree_idx - 1, 1)
|
||||
for i in range(draft_token_num - 1):
|
||||
tl.store(tree_mask_ptr + token_tree_idx + i, 0)
|
||||
|
||||
if draft_tokenx > 0:
|
||||
# Build tree path for draft_tokenx > 0
|
||||
cur_position = draft_tokenx - 1
|
||||
position = 0
|
||||
should_continue = 1
|
||||
|
||||
for _ in range(depth):
|
||||
if should_continue:
|
||||
position += 1
|
||||
tl.store(tree_mask_ptr + token_tree_idx + cur_position, 1)
|
||||
|
||||
parent_tb_idx = (
|
||||
tl.load(
|
||||
selected_index_ptr
|
||||
+ batch_idx * selected_index_stride
|
||||
+ cur_position
|
||||
)
|
||||
// topk
|
||||
)
|
||||
if parent_tb_idx == 0:
|
||||
should_continue = 0
|
||||
else:
|
||||
parent_token_idx = tl.load(
|
||||
parent_list_ptr
|
||||
+ batch_idx * parent_list_stride
|
||||
+ parent_tb_idx
|
||||
)
|
||||
|
||||
# Find cur_position for next iteration
|
||||
found = 0
|
||||
for cp in range(draft_token_num - 1):
|
||||
if found == 0:
|
||||
if (
|
||||
tl.load(
|
||||
selected_index_ptr
|
||||
+ batch_idx * selected_index_stride
|
||||
+ cp
|
||||
)
|
||||
== parent_token_idx
|
||||
):
|
||||
cur_position = cp
|
||||
found = 1
|
||||
if found == 0:
|
||||
should_continue = 0
|
||||
|
||||
tl.store(
|
||||
positions_ptr + batch_idx * draft_token_num + draft_tokenx,
|
||||
position + seq_len,
|
||||
)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def verify_tree_greedy_kernel_triton(
|
||||
predicts_ptr,
|
||||
accept_index_ptr,
|
||||
accept_token_num_ptr,
|
||||
candidates_ptr,
|
||||
retrieve_index_ptr,
|
||||
retrieve_next_token_ptr,
|
||||
retrieve_next_sibling_ptr,
|
||||
target_predict_ptr,
|
||||
batch_size: tl.constexpr,
|
||||
num_speculative_tokens: tl.constexpr,
|
||||
num_draft_tokens: tl.constexpr,
|
||||
):
|
||||
"""
|
||||
Triton kernel for verifying EAGLE tree in greedy mode.
|
||||
Each program handles one batch item.
|
||||
"""
|
||||
bx = tl.program_id(0)
|
||||
|
||||
# Initialize
|
||||
last_accept_retrieve_idx = tl.load(retrieve_index_ptr + bx * num_draft_tokens)
|
||||
tl.store(accept_index_ptr + bx * num_speculative_tokens, last_accept_retrieve_idx)
|
||||
# Cast to match dtype of loaded tensors to avoid type inconsistency
|
||||
num_accept_tokens = tl.cast(0, last_accept_retrieve_idx.dtype)
|
||||
cur_index = tl.cast(0, last_accept_retrieve_idx.dtype)
|
||||
|
||||
# Tree traversal loop
|
||||
should_continue = 1
|
||||
for j in range(1, num_speculative_tokens):
|
||||
if should_continue: # Early exit guard
|
||||
cur_index = tl.load(
|
||||
retrieve_next_token_ptr + bx * num_draft_tokens + cur_index
|
||||
)
|
||||
|
||||
# Load target token once per level (before sibling search)
|
||||
# last_accept_retrieve_idx is constant during sibling traversal
|
||||
target_row = last_accept_retrieve_idx // num_draft_tokens
|
||||
target_col = last_accept_retrieve_idx % num_draft_tokens
|
||||
target_token = tl.load(
|
||||
target_predict_ptr + target_row * num_draft_tokens + target_col
|
||||
)
|
||||
|
||||
# Traverse siblings
|
||||
found_match = 0
|
||||
for _ in range(num_draft_tokens): # Max iterations = num_draft_tokens
|
||||
if found_match == 0: # Early exit guard
|
||||
# Check if we've reached end of sibling list
|
||||
is_valid = cur_index != -1
|
||||
|
||||
# Use masked loads with safe address (0 when invalid)
|
||||
safe_cur_index = (
|
||||
cur_index * is_valid
|
||||
) # 0 if invalid, cur_index if valid
|
||||
safe_index = bx * num_draft_tokens + safe_cur_index
|
||||
|
||||
# Load draft token info (loads from index 0 when invalid, but we won't use it)
|
||||
draft_index = tl.load(retrieve_index_ptr + safe_index)
|
||||
draft_token = tl.load(candidates_ptr + safe_index)
|
||||
|
||||
# Check for token match (only valid when is_valid is True)
|
||||
token_match = is_valid & (draft_token == target_token)
|
||||
|
||||
# Accept token using predicated stores (only write if matched)
|
||||
tl.store(
|
||||
predicts_ptr + last_accept_retrieve_idx,
|
||||
target_token,
|
||||
mask=token_match,
|
||||
)
|
||||
next_num_accept_tokens = num_accept_tokens + 1
|
||||
tl.store(
|
||||
accept_index_ptr
|
||||
+ bx * num_speculative_tokens
|
||||
+ next_num_accept_tokens,
|
||||
draft_index,
|
||||
mask=token_match,
|
||||
)
|
||||
|
||||
num_accept_tokens = num_accept_tokens + token_match
|
||||
last_accept_retrieve_idx = (
|
||||
token_match * draft_index
|
||||
+ (~token_match) * last_accept_retrieve_idx
|
||||
)
|
||||
found_match = token_match * 1 + (~is_valid) * (-1)
|
||||
|
||||
# Masked load: only load next sibling when no match (hardware predication)
|
||||
# When matched: returns cur_index (other); when not matched: loads sibling
|
||||
cur_index = tl.load(
|
||||
retrieve_next_sibling_ptr + safe_index,
|
||||
mask=~token_match
|
||||
& is_valid, # Only load when valid and NOT matched
|
||||
other=cur_index, # Keep cur_index when matched or invalid
|
||||
)
|
||||
|
||||
if found_match != 1:
|
||||
should_continue = 0
|
||||
|
||||
# Store final results
|
||||
tl.store(accept_token_num_ptr + bx, num_accept_tokens)
|
||||
|
||||
target_row = last_accept_retrieve_idx // num_draft_tokens
|
||||
target_col = last_accept_retrieve_idx % num_draft_tokens
|
||||
final_target = tl.load(
|
||||
target_predict_ptr + target_row * num_draft_tokens + target_col
|
||||
)
|
||||
tl.store(predicts_ptr + last_accept_retrieve_idx, final_target)
|
||||
Reference in New Issue
Block a user