[SPEC_V2] Enable cudagraph draft_extend for trtllm_mla_backend and Acclen Fix for DP under cudagraph mode (#16974)

This commit is contained in:
YAMY
2026-01-18 15:56:21 +08:00
committed by GitHub
parent f78201f3a9
commit a45e0e5df4
2 changed files with 22 additions and 3 deletions
@@ -253,6 +253,12 @@ class EAGLEDraftExtendCudaGraphRunner:
: bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens : bs if self.forward_mode == ForwardMode.DRAFT_EXTEND else num_tokens
] ]
# V1 (DRAFT_EXTEND): pruned_states = bs (last token per seq)
# V2 (DRAFT_EXTEND_V2): pruned_states = num_tokens (all tokens)
num_tokens_for_logprob = (
num_tokens if self.forward_mode.is_draft_extend_v2() else bs
)
if self.require_mlp_tp_gather: if self.require_mlp_tp_gather:
self.global_num_tokens_gpu.copy_( self.global_num_tokens_gpu.copy_(
torch.tensor( torch.tensor(
@@ -263,7 +269,7 @@ class EAGLEDraftExtendCudaGraphRunner:
) )
self.global_num_tokens_for_logprob_gpu.copy_( self.global_num_tokens_for_logprob_gpu.copy_(
torch.tensor( torch.tensor(
[bs] * self.dp_size, [num_tokens_for_logprob] * self.dp_size,
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=self.input_ids.device,
) )
@@ -279,7 +285,7 @@ class EAGLEDraftExtendCudaGraphRunner:
) )
self.global_num_tokens_for_logprob_gpu.copy_( self.global_num_tokens_for_logprob_gpu.copy_(
torch.tensor( torch.tensor(
[bs], [num_tokens_for_logprob],
dtype=torch.int32, dtype=torch.int32,
device=self.input_ids.device, device=self.input_ids.device,
) )
@@ -419,7 +425,13 @@ class EAGLEDraftExtendCudaGraphRunner:
# TODO(ch-wan): support num_token_non_padded # TODO(ch-wan): support num_token_non_padded
if self.require_gathered_buffer: if self.require_gathered_buffer:
self.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs) self.global_num_tokens_gpu.fill_(bs * self.num_tokens_per_bs)
self.global_num_tokens_for_logprob_gpu.fill_(bs) # V1: pruned_states = bs; V2: pruned_states = num_tokens
if self.forward_mode.is_draft_extend_v2():
self.global_num_tokens_for_logprob_gpu.fill_(
bs * self.num_tokens_per_bs
)
else:
self.global_num_tokens_for_logprob_gpu.fill_(bs)
if forward_batch.seq_lens_cpu is not None: if forward_batch.seq_lens_cpu is not None:
if bs != raw_bs: if bs != raw_bs:
@@ -13,6 +13,9 @@ from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner i
EAGLEDraftNpuGraphRunner, EAGLEDraftNpuGraphRunner,
) )
from sglang.srt.layers.attention.triton_backend import TritonMultiStepDraftBackend from sglang.srt.layers.attention.triton_backend import TritonMultiStepDraftBackend
from sglang.srt.layers.attention.trtllm_mla_backend import (
TRTLLMMLAMultiStepDraftBackend,
)
from sglang.srt.layers.dp_attention import get_attention_tp_group from sglang.srt.layers.dp_attention import get_attention_tp_group
from sglang.srt.layers.moe.utils import ( from sglang.srt.layers.moe.utils import (
speculative_moe_a2a_backend_context, speculative_moe_a2a_backend_context,
@@ -277,6 +280,10 @@ class EagleDraftWorker(BaseDraftWorker):
_is_cuda _is_cuda
and isinstance(self.draft_attn_backend, TritonMultiStepDraftBackend) and isinstance(self.draft_attn_backend, TritonMultiStepDraftBackend)
) )
or (
_is_cuda
and isinstance(self.draft_attn_backend, TRTLLMMLAMultiStepDraftBackend)
)
): ):
tic = time.perf_counter() tic = time.perf_counter()
before_mem = get_available_gpu_memory(self.device, self.gpu_id) before_mem = get_available_gpu_memory(self.device, self.gpu_id)