From a76b74cbe0114c58fbc911cb1180436c09135711 Mon Sep 17 00:00:00 2001 From: ming_wang <68357922+sigama-w@users.noreply.github.com> Date: Sun, 26 Jul 2026 20:42:14 +0800 Subject: [PATCH] add fill_draft_extend_prepare_buffers_native for NPU (#32427) --- .../ops/speculative/multi_layer_eagle.py | 66 +++++++++++++++++++ .../sglang/srt/disaggregation/ascend/conn.py | 3 +- ...er_eagle_draft_extend_cuda_graph_runner.py | 14 +++- .../speculative/multi_layer_eagle_utils.py | 2 + 4 files changed, 82 insertions(+), 3 deletions(-) diff --git a/python/sglang/kernels/ops/speculative/multi_layer_eagle.py b/python/sglang/kernels/ops/speculative/multi_layer_eagle.py index 3d3ef92f3..e088e52c7 100644 --- a/python/sglang/kernels/ops/speculative/multi_layer_eagle.py +++ b/python/sglang/kernels/ops/speculative/multi_layer_eagle.py @@ -728,3 +728,69 @@ def fill_draft_extend_prepare_buffers_triton( BLOCK_HIDDEN=BLOCK_HIDDEN, GLOBAL_BLOCK=global_block, ) + + +def fill_draft_extend_prepare_buffers_native( + input_ids, + positions, + out_cache_loc, + src_input_ids, + src_positions, + src_out_cache_loc, + seq_lens, + req_pool_indices, + num_correct_drafts, + num_accept_tokens, + select_index, + temperatures, + src_seq_lens, + src_req_pool_indices, + src_num_correct_drafts, + src_num_accept_tokens, + src_temperatures, + hidden_states, + src_hidden_states, + global_num_tokens, + global_num_tokens_for_logprob, + raw_bs, + bs, + num_tokens_per_bs, + num_front_tokens, + seq_len_fill_value, +): + """Native PyTorch implementation of fill_draft_extend_prepare_buffers. + + Used on NPU where the triton mega-kernel triggers CCU errors. + Mirrors the triton kernel's semantics exactly. + """ + num_tokens = src_input_ids.shape[0] + + # Token buffers: copy real rows, leave padding untouched. + input_ids[:num_tokens].copy_(src_input_ids) + positions[:num_tokens].copy_(src_positions) + out_cache_loc[:num_tokens].copy_(src_out_cache_loc) + + # Per-request buffers. + seq_lens[:bs].fill_(seq_len_fill_value) + seq_lens[:raw_bs].copy_(src_seq_lens) + req_pool_indices[:raw_bs].copy_(src_req_pool_indices) + + num_correct_drafts[:raw_bs].copy_(src_num_correct_drafts) + num_accept_tokens[:bs].fill_(-1) + num_accept_tokens[:raw_bs].copy_(src_num_accept_tokens) + + # select_index = i * num_tokens_per_bs + num_front_tokens + num_correct_drafts + idx = torch.arange(bs, device=select_index.device, dtype=torch.int64) + select_index[:bs] = idx * num_tokens_per_bs + num_front_tokens + select_index[:bs] += num_correct_drafts[:bs].to(torch.int64) + + if temperatures is not None: + temperatures[:bs].fill_(1.0) + temperatures[:raw_bs].copy_(src_temperatures) + + if global_num_tokens is not None: + global_num_tokens.fill_(bs * num_tokens_per_bs) + global_num_tokens_for_logprob.fill_(bs * num_tokens_per_bs) + + if src_hidden_states is not None: + hidden_states[:num_tokens].copy_(src_hidden_states) diff --git a/python/sglang/srt/disaggregation/ascend/conn.py b/python/sglang/srt/disaggregation/ascend/conn.py index 90b020846..40251c587 100644 --- a/python/sglang/srt/disaggregation/ascend/conn.py +++ b/python/sglang/srt/disaggregation/ascend/conn.py @@ -1,6 +1,6 @@ import concurrent.futures import logging -from typing import List, Tuple +from typing import List, Optional, Tuple import numpy as np import numpy.typing as npt @@ -78,6 +78,7 @@ class AscendKVManager(MooncakeKVManager): dst_kv_ptrs: list[int], dst_kv_indices: npt.NDArray[np.int32], executor: concurrent.futures.ThreadPoolExecutor, + dst_layer_ids: Optional[List[int]] = None, ): # Group by indices prefill_kv_blocks, dst_kv_blocks = group_concurrent_contiguous( diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index db21f9e06..f4bd3fb16 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -63,7 +63,9 @@ from sglang.srt.runtime_context import get_flags from sglang.srt.speculative.eagle_info import EagleDraftExtendInput from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim from sglang.srt.speculative.multi_layer_eagle_utils import ( - fill_draft_extend_prepare_buffers_triton, + fill_draft_extend_prepare_buffers_triton as fill_draft_extend_prepare_buffers, +) +from sglang.srt.speculative.multi_layer_eagle_utils import ( rotate_input_ids, wide_row_softmax_triton, ) @@ -74,12 +76,20 @@ from sglang.srt.speculative.spec_utils import ( ) from sglang.srt.utils import ( get_available_gpu_memory, + is_npu, require_attn_tp_gather, require_gathered_buffer, require_mlp_sync, require_mlp_tp_gather, ) +if is_npu(): + from sglang.srt.speculative.multi_layer_eagle_utils import ( + fill_draft_extend_prepare_buffers_native, + ) + + fill_draft_extend_prepare_buffers = fill_draft_extend_prepare_buffers_native + if TYPE_CHECKING: from sglang.srt.speculative.multi_layer_eagle_worker_v2 import ( MultiLayerEagleDraftWorker, @@ -713,7 +723,7 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner: else: bs = self.get_runner(0)._pad_to_bucket(raw_bs, self.capture_bs) - fill_draft_extend_prepare_buffers_triton( + fill_draft_extend_prepare_buffers( buffers.input_ids, buffers.positions, buffers.out_cache_loc, diff --git a/python/sglang/srt/speculative/multi_layer_eagle_utils.py b/python/sglang/srt/speculative/multi_layer_eagle_utils.py index e44735d1a..5c93c5fa6 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_utils.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_utils.py @@ -14,6 +14,7 @@ from sglang.kernels.ops.speculative.multi_layer_eagle import ( compute_widened_draft_extend_locs_positions_triton, + fill_draft_extend_prepare_buffers_native, fill_draft_extend_prepare_buffers_triton, fill_widened_draft_extend_inputs_triton, rotate_input_ids, @@ -53,6 +54,7 @@ def compute_widened_draft_extend_locs_positions( __all__ = [ "boundary_kv_fix_enabled", "compute_widened_draft_extend_locs_positions", + "fill_draft_extend_prepare_buffers_native", "fill_draft_extend_prepare_buffers_triton", "fill_widened_draft_extend_inputs_triton", "rotate_input_ids",