[NPU] Fix mrope_position computation in Eagle Worker v2 with PlanStream (#23423)
This commit is contained in:
@@ -28,6 +28,7 @@ import torch
|
||||
import sglang
|
||||
from sglang.srt.configs.model_config import AttentionArch, is_deepseek_nsa
|
||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.model_executor.cuda_graph_runner import CudaGraphRunner
|
||||
from sglang.srt.utils import (
|
||||
empty_context,
|
||||
@@ -173,6 +174,13 @@ class NPUGraphRunner(CudaGraphRunner):
|
||||
# In speculative decoding, these two fields are still needed.
|
||||
self.buffers.input_ids[: self.raw_num_token].copy_(forward_batch.input_ids)
|
||||
self.buffers.positions[: self.raw_num_token].copy_(forward_batch.positions)
|
||||
if (
|
||||
envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.get()
|
||||
and forward_batch.mrope_positions is not None
|
||||
):
|
||||
self.buffers.mrope_positions[:, : self.raw_num_token].copy_(
|
||||
forward_batch.mrope_positions
|
||||
)
|
||||
|
||||
self.update_attr_name = self._get_update_attr_name()
|
||||
self.update_attr_type = self._get_update_attr_type()
|
||||
|
||||
@@ -587,7 +587,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
ret.spec_info is not None
|
||||
and getattr(ret.spec_info, "positions", None) is not None
|
||||
):
|
||||
ret._compute_spec_mrope_positions(model_runner, batch)
|
||||
ret.compute_spec_mrope_positions(model_runner, batch)
|
||||
else:
|
||||
ret._compute_mrope_positions(model_runner, batch)
|
||||
|
||||
@@ -696,7 +696,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
req_lens=req_lens,
|
||||
)
|
||||
|
||||
def _compute_spec_mrope_positions(
|
||||
def compute_spec_mrope_positions(
|
||||
self, model_runner: ModelRunner, batch: ModelWorkerBatch
|
||||
):
|
||||
# TODO support batched deltas
|
||||
|
||||
@@ -969,6 +969,18 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
||||
torch.get_device_module(self.device).current_stream().wait_stream(
|
||||
self.plan_stream
|
||||
)
|
||||
if (
|
||||
_is_npu
|
||||
and self._target_worker.model_runner.model_is_mrope
|
||||
and batch.spec_info is not None
|
||||
and getattr(batch.spec_info, "positions", None) is not None
|
||||
and not batch.forward_mode.is_idle()
|
||||
):
|
||||
# mrope_position depends on draft output in default stream and is computed in plan stream,
|
||||
# causing errors. Compute it here for correct values.
|
||||
verify_forward_batch.compute_spec_mrope_positions(
|
||||
self._target_worker.model_runner, batch
|
||||
)
|
||||
|
||||
# Some values such as custom_mask and position depend on the output of draft,
|
||||
# so the previous plan step used the wrong values. Here, we need to run the related
|
||||
|
||||
Reference in New Issue
Block a user