Support speculative decoding on CPU (#27862)

Co-authored-by: Valentine233 <xuan.liao@intel.com>
This commit is contained in:
Haotong Zou
2026-07-09 10:27:09 +08:00
committed by GitHub
co-authored by Valentine233
parent 177c048c68
commit 3b43df5b6d
36 changed files with 3499 additions and 138 deletions
+10
View File
@@ -76,6 +76,7 @@ else:
max_pooling_1d_varlen,
)
from sgl_kernel.kvcacheio import (
copy_all_layer_kv_cache_cpu,
transfer_kv_all_layer,
transfer_kv_all_layer_mla,
transfer_kv_per_layer,
@@ -113,11 +114,20 @@ else:
top_p_renorm_prob,
)
from sgl_kernel.speculative import (
assign_draft_cache_locs_contiguous_cpu,
assign_extend_cache_locs_cpu,
assign_req_to_token_pool_cpu,
build_draft_decode_metadata_cpu,
build_tree_kernel_efficient,
build_tree_kernel_efficient_cpu,
fill_accept_out_cache_loc_cpu,
fill_bonus_tokens_cpu,
reconstruct_indices_from_tree_mask,
rotate_input_ids_cpu,
segment_packbits,
tree_speculative_sampling_target_only,
verify_tree_greedy,
verify_tree_greedy_cpu,
)
from sgl_kernel.top_k import (
fast_topk,
+14
View File
@@ -305,3 +305,17 @@ def transfer_kv_all_layer_mla_lf_pf(
block_quota,
num_warps_per_block,
)
def copy_all_layer_kv_cache_cpu(
data_ptrs: torch.Tensor,
strides: torch.Tensor,
tgt_loc: torch.Tensor,
src_loc: torch.Tensor,
):
torch.ops.sgl_kernel.copy_all_layer_kv_cache_cpu(
data_ptrs,
strides,
tgt_loc,
src_loc,
)
+192 -10
View File
@@ -1,3 +1,5 @@
from typing import Optional
import torch
@@ -97,16 +99,28 @@ def reconstruct_indices_from_tree_mask(
batch_size: int,
draft_token_num: int,
) -> None:
torch.ops.sgl_kernel.reconstruct_indices_from_tree_mask.default(
tree_mask,
verified_seq_len,
positions,
retrive_index,
retrive_next_token,
retrive_next_sibling,
batch_size,
draft_token_num,
)
if tree_mask.is_cpu:
torch.ops.sgl_kernel.reconstruct_indices_from_tree_mask_cpu(
tree_mask,
verified_seq_len,
positions,
retrive_index,
retrive_next_token,
retrive_next_sibling,
batch_size,
draft_token_num,
)
else:
torch.ops.sgl_kernel.reconstruct_indices_from_tree_mask.default(
tree_mask,
verified_seq_len,
positions,
retrive_index,
retrive_next_token,
retrive_next_sibling,
batch_size,
draft_token_num,
)
def segment_packbits(
@@ -124,3 +138,171 @@ def segment_packbits(
batch_size,
torch.cuda.current_stream().cuda_stream,
)
def verify_tree_greedy_cpu(
predicts: torch.Tensor, # mutable
accept_index: torch.Tensor, # mutable
accept_token_num: torch.Tensor, # mutable
candidates: torch.Tensor,
retrive_index: torch.Tensor,
retrive_next_token: torch.Tensor,
retrive_next_sibling: torch.Tensor,
target_predict: torch.Tensor,
) -> None:
torch.ops.sgl_kernel.verify_tree_greedy_cpu(
predicts,
accept_index,
accept_token_num,
candidates,
retrive_index,
retrive_next_token,
retrive_next_sibling,
target_predict,
)
def build_tree_kernel_efficient_cpu(
parent_list: torch.Tensor,
selected_index: torch.Tensor,
verified_seq_len: torch.Tensor,
tree_mask: torch.Tensor,
positions: torch.Tensor,
retrive_index: torch.Tensor,
retrive_next_token: torch.Tensor,
retrive_next_sibling: torch.Tensor,
topk: int,
depth: int,
draft_token_num: int,
tree_mask_mode: int,
) -> None:
torch.ops.sgl_kernel.build_tree_kernel_efficient_cpu(
parent_list,
selected_index,
verified_seq_len,
tree_mask,
positions,
retrive_index,
retrive_next_token,
retrive_next_sibling,
topk,
depth,
draft_token_num,
tree_mask_mode,
)
def assign_req_to_token_pool_cpu(
req_pool_indices: torch.Tensor,
req_to_token: torch.Tensor,
start_offset: torch.Tensor,
end_offset: torch.Tensor,
out_cache_loc: torch.Tensor,
pool_len: int,
) -> None:
torch.ops.sgl_kernel.assign_req_to_token_pool_cpu(
req_pool_indices,
req_to_token,
start_offset,
end_offset,
out_cache_loc,
pool_len,
)
def build_draft_decode_metadata_cpu(
req_to_token: torch.Tensor,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
topk: int,
num_steps: int,
pool_len: int,
) -> torch.Tensor:
return torch.ops.sgl_kernel.build_draft_decode_metadata_cpu(
req_to_token,
req_pool_indices,
seq_lens,
topk,
num_steps,
pool_len,
)
def fill_bonus_tokens_cpu(
accept_tokens: torch.Tensor,
accept_lens: torch.Tensor,
bonus_tokens: torch.Tensor,
accept_stride: int,
) -> None:
torch.ops.sgl_kernel.fill_bonus_tokens_cpu(
accept_tokens,
accept_lens,
bonus_tokens,
accept_stride,
)
def fill_accept_out_cache_loc_cpu(
accept_index: torch.Tensor,
out_cache_loc: torch.Tensor,
accept_out_cache_loc: torch.Tensor, # mutable
) -> None:
torch.ops.sgl_kernel.fill_accept_out_cache_loc_cpu(
accept_index,
out_cache_loc,
accept_out_cache_loc,
)
def assign_draft_cache_locs_contiguous_cpu(
req_pool_indices: torch.Tensor,
req_to_token: torch.Tensor,
seq_lens: torch.Tensor,
out_cache_loc: torch.Tensor,
pool_len: int,
topk: int,
num_steps: int,
) -> None:
torch.ops.sgl_kernel.assign_draft_cache_locs_contiguous_cpu(
req_pool_indices,
req_to_token,
seq_lens,
out_cache_loc,
pool_len,
topk,
num_steps,
)
def assign_extend_cache_locs_cpu(
req_pool_indices: torch.Tensor,
req_to_token: torch.Tensor,
start_offset: torch.Tensor,
end_offset: torch.Tensor,
out_cache_loc: torch.Tensor,
pool_len: int,
) -> None:
torch.ops.sgl_kernel.assign_extend_cache_locs_cpu(
req_pool_indices,
req_to_token,
start_offset,
end_offset,
out_cache_loc,
pool_len,
)
def rotate_input_ids_cpu(
input_ids: torch.Tensor,
extend_start_loc: torch.Tensor,
extend_seq_lens: torch.Tensor,
topk_index: torch.Tensor,
select_index: Optional[torch.Tensor] = None,
) -> None:
torch.ops.sgl_kernel.rotate_input_ids_cpu(
input_ids,
extend_start_loc,
extend_seq_lens,
topk_index,
select_index,
)