[nemetron/mtp] fix nemotron mtp (#16275)

This commit is contained in:
Hanming Lu
2026-01-09 07:22:21 -08:00
committed by GitHub
parent 8eeffbe9aa
commit 5dcff94791
3 changed files with 33 additions and 3 deletions

View File

@@ -534,6 +534,9 @@ class MambaMixer2(torch.nn.Module):
layer_cache, MambaPool.SpeculativeState
), "layer_cache must be SpeculativeState for speculative decoding"
draft_token_num = metadata.draft_token_num
self.intermediate_state_indices = torch.arange(
num_decodes, dtype=torch.int32, device=state_indices_tensor_d.device
)
# Reshape for batch processing
hidden_states_B_C_d_reshaped = hidden_states_B_C_d.view(
@@ -548,6 +551,7 @@ class MambaMixer2(torch.nn.Module):
self.activation,
conv_state_indices=state_indices_tensor_d[:num_decodes],
intermediate_conv_window=layer_cache.intermediate_conv_window[0],
intermediate_state_indices=self.intermediate_state_indices,
retrieve_next_token=metadata.retrieve_next_token,
retrieve_next_sibling=metadata.retrieve_next_sibling,
retrieve_parent_token=metadata.retrieve_parent_token,
@@ -621,6 +625,7 @@ class MambaMixer2(torch.nn.Module):
intermediate_states_buffer=layer_cache.intermediate_ssm,
cache_steps=draft_token_num,
retrieve_parent_token=metadata.retrieve_parent_token,
intermediate_state_indices=self.intermediate_state_indices,
)
else:
selective_state_update(

View File

@@ -56,6 +56,14 @@ else:
is not None
}
)
@triton.heuristics(
{
"HAS_INTERMEDIATE_STATE_INDICES": lambda args: args[
"intermediate_state_indices_ptr"
]
is not None
}
)
@triton.jit(do_not_specialize=["T"])
def _selective_scan_update_kernel(
# Pointers to matrices
@@ -74,6 +82,7 @@ def _selective_scan_update_kernel(
intermediate_states_buffer,
cache_steps,
retrieve_parent_token_ptr,
intermediate_state_indices_ptr,
# Matrix dimensions
batch,
T,
@@ -130,6 +139,7 @@ def _selective_scan_update_kernel(
DISABLE_STATE_UPDATE: tl.constexpr,
CACHE_INTERMEDIATE_STATES: tl.constexpr,
HAS_EAGLE_TREE_CUSTOM_ATTN_MASK: tl.constexpr,
HAS_INTERMEDIATE_STATE_INDICES: tl.constexpr,
BLOCK_SIZE_DSTATE: tl.constexpr,
):
pid_m = tl.program_id(axis=0)
@@ -177,7 +187,12 @@ def _selective_scan_update_kernel(
cache_idx = -1
if CACHE_INTERMEDIATE_STATES:
if HAS_STATE_BATCH_INDICES:
if HAS_INTERMEDIATE_STATE_INDICES:
intermediate_state_idx = tl.load(intermediate_state_indices_ptr + pid_b).to(
tl.int64
)
cache_idx = intermediate_state_idx
elif HAS_STATE_BATCH_INDICES:
cache_idx = state_batch_idx
else:
cache_idx = pid_b
@@ -250,7 +265,7 @@ def _selective_scan_update_kernel(
if state_batch_idx != pad_slot_id:
cache_ptr_base = (
intermediate_states_buffer
+ state_batch_idx * cache_steps * nheads * dim * dstate
+ cache_idx * cache_steps * nheads * dim * dstate
+ current_step_idx * nheads * dim * dstate
+ pid_h * dim * dstate
)
@@ -300,6 +315,7 @@ def selective_state_update(
intermediate_states_buffer=None,
cache_steps=None,
retrieve_parent_token=None,
intermediate_state_indices=None,
):
"""
Argument:
@@ -324,6 +340,8 @@ def selective_state_update(
intermediate_states_buffer: Buffer to cache intermediate states
cache_steps: Total number of steps in the buffer
retrieve_parent_token: (batch, T) tensor of parent token indices for EAGLE tree attention
intermediate_state_indices: (batch,) tensor of indices for intermediate_states_buffer operations.
If provided, uses these indices instead of state_batch_indices for the buffer.
"""
if state.dim() == 3:
state = state.unsqueeze(1)
@@ -426,6 +444,7 @@ def selective_state_update(
intermediate_states_buffer,
cache_steps if cache_steps is not None else 0,
retrieve_parent_token,
intermediate_state_indices,
batch,
T,
nheads,

View File

@@ -1376,7 +1376,13 @@ class ServerArgs:
else:
self.quantization = model_config.quantization
self.moe_runner_backend = "flashinfer_cutlass"
if not self.disable_radix_cache:
if not self.disable_radix_cache and self.speculative_algorithm is not None:
logger.warning(
"Disabling radix cache since speculative decoding for NemotronHForCausalLM is not supported with radix cache yet."
)
self.disable_radix_cache = True
elif not self.disable_radix_cache:
logger.warning(
"Disabling overlap schedule since MambaRadixCache is not compatible with "
"overlap schedule currently, try to use --disable-radix-cache if overlap schedule is necessary"