Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
472 changes: 326 additions & 146 deletions tensorrt_llm/_torch/models/modeling_speculative.py

Large diffs are not rendered by default.

6 changes: 3 additions & 3 deletions tensorrt_llm/_torch/modules/mamba/gdn_mixer.py
Original file line number Diff line number Diff line change
Expand Up @@ -466,7 +466,7 @@ def forward_decode(
num_decodes = kwargs["num_decodes"]

if is_target_verify:
draft_token_num = spec_metadata.max_draft_len + 1
draft_token_num = spec_metadata.max_total_draft_tokens + 1
assert num_decodes > 0
assert mixed_qkv.shape[0] == num_decodes * draft_token_num
assert a.shape[0] == num_decodes * draft_token_num
Expand Down Expand Up @@ -634,7 +634,7 @@ def forward_extend(
)

if is_target_verify:
draft_token_num = spec_metadata.max_draft_len + 1
draft_token_num = spec_metadata.max_total_draft_tokens + 1
assert num_decodes > 0
assert mixed_qkv_d.shape[0] == num_decodes * draft_token_num
assert a_d.shape[0] == num_decodes * draft_token_num
Expand Down Expand Up @@ -738,7 +738,7 @@ def forward_extend(
last_recurrent_state = last_recurrent_state.to(ssm_states.dtype, copy=False)
ssm_states[state_indices_p] = last_recurrent_state

draft_token_num = spec_metadata.max_draft_len + 1
draft_token_num = spec_metadata.max_total_draft_tokens + 1
query_d = query[:, num_prefill_tokens:, :, :].reshape(
num_decodes, draft_token_num, self.num_k_heads // self.attn_tp_size, self.head_k_dim
)
Expand Down
2 changes: 1 addition & 1 deletion tensorrt_llm/_torch/pyexecutor/mamba_cache_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -868,7 +868,7 @@ def __init__(
mamba_ssm_cache_dtype,
mamba_layer_mask,
execution_stream,
speculative_num_draft_tokens=(spec_config.max_draft_len
speculative_num_draft_tokens=(spec_config.tokens_per_gen_step - 1
if spec_config is not None else None),
model_type=model_type,
use_replay_state_update=use_replay_state_update,
Expand Down
35 changes: 26 additions & 9 deletions tensorrt_llm/_torch/pyexecutor/model_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -420,8 +420,12 @@ def __init__(
self.without_logits = self.spec_config.spec_dec_mode.without_logits(
) or self.model_is_wrapped
self.max_total_draft_tokens = spec_config.tokens_per_gen_step - 1
# PARD/DFlash use 2K tokens per gen request (K accepted + K masks), so
# their per-request draft buffer width is 2K-1 = max_total_draft_tokens.
# Parallel-draft modes (PARD, DFlash) size their per-request draft
# buffer by tokens_per_gen_step - 1 so the engine reserves exactly
# one slot per draft token the target will verify. PARD still uses
# 2K tokens per gen req (K drafts + K mask fillers); DFlash was
# reduced to K+1 (K drafts + 1 bonus) - the spec config's
# tokens_per_gen_step carries the per-algorithm width.
if spec_config.spec_dec_mode.is_parallel_draft():
self.max_draft_len = self.max_total_draft_tokens
else:
Expand Down Expand Up @@ -1636,8 +1640,14 @@ def _preprocess_inputs(self, inputs: Dict[str, Any]):
'attn_metadata'].num_chunked_ctx_requests
previous_batch_tokens = inputs['input_ids'].shape[
0] - num_ctx_tokens
inputs['position_ids'][0, num_ctx_tokens:] += (
self.previous_pos_id_offsets_cuda[:previous_batch_tokens])
if inputs['position_ids'].ndim == 3: # mrope: [3, 1, N]
inputs['position_ids'][:, :, num_ctx_tokens:] += (
self.
previous_pos_id_offsets_cuda[:previous_batch_tokens])
else:
inputs['position_ids'][0, num_ctx_tokens:] += (
self.
previous_pos_id_offsets_cuda[:previous_batch_tokens])
if hasattr(inputs['attn_metadata'], 'kv_lens_cuda'):
if num_ctx_requests >= num_chunked_ctx_requests and num_chunked_ctx_requests > 0:
# The generation requests with draft_tokens are treated as chunked context requests when extend_ctx returns True.
Expand Down Expand Up @@ -1677,8 +1687,14 @@ def _postprocess_inputs(self, inputs: Dict[str, Any]):
'attn_metadata'].num_chunked_ctx_requests
previous_batch_tokens = inputs['input_ids'].shape[
0] - num_ctx_tokens
inputs['position_ids'][0, num_ctx_tokens:] -= (
self.previous_pos_id_offsets_cuda[:previous_batch_tokens])
if inputs['position_ids'].ndim == 3: # mrope: [3, 1, N]
inputs['position_ids'][:, :, num_ctx_tokens:] -= (
self.
previous_pos_id_offsets_cuda[:previous_batch_tokens])
else:
inputs['position_ids'][0, num_ctx_tokens:] -= (
self.
previous_pos_id_offsets_cuda[:previous_batch_tokens])
# Only TrtllmAttentionMetadata has kv_lens_cuda.
if isinstance(inputs['attn_metadata'], TrtllmAttentionMetadata):
if num_ctx_requests >= num_chunked_ctx_requests and num_chunked_ctx_requests > 0:
Expand Down Expand Up @@ -3886,9 +3902,10 @@ def forward(self,
# to spec_metadata so downstream code (eagle3, interface, trtllm) can read it.
spec_metadata.runtime_draft_len = self.runtime_draft_len

# PARD/DFlash have 2K tokens per gen request, not K+1. Pass 2K-1
# so generation_lengths = 2K and the XQA kernel computes
# the correct past_kv_len.
# Parallel-draft modes advertise a per-gen-step width via
# tokens_per_gen_step (PARD: 2K, DFlash: K+1). Pass
# (tokens_per_gen_step - 1) so generation_lengths = tokens_per_gen_step
# and the XQA kernel computes the correct past_kv_len.
if spec_metadata.spec_dec_mode.is_parallel_draft():
sd_max_draft_len = self.original_max_total_draft_tokens
sd_max_total = self.original_max_total_draft_tokens
Expand Down
Loading
Loading