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

This commit is contained in:
Hanming Lu
2026-01-09 23:22:21 +08:00
committed by GitHub
parent 8eeffbe9aa
commit 5dcff94791
3 changed files with 33 additions and 3 deletions
@@ -534,6 +534,9 @@ class MambaMixer2(torch.nn.Module):
layer_cache, MambaPool.SpeculativeState layer_cache, MambaPool.SpeculativeState
), "layer_cache must be SpeculativeState for speculative decoding" ), "layer_cache must be SpeculativeState for speculative decoding"
draft_token_num = metadata.draft_token_num 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 # Reshape for batch processing
hidden_states_B_C_d_reshaped = hidden_states_B_C_d.view( hidden_states_B_C_d_reshaped = hidden_states_B_C_d.view(
@@ -548,6 +551,7 @@ class MambaMixer2(torch.nn.Module):
self.activation, self.activation,
conv_state_indices=state_indices_tensor_d[:num_decodes], conv_state_indices=state_indices_tensor_d[:num_decodes],
intermediate_conv_window=layer_cache.intermediate_conv_window[0], 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_token=metadata.retrieve_next_token,
retrieve_next_sibling=metadata.retrieve_next_sibling, retrieve_next_sibling=metadata.retrieve_next_sibling,
retrieve_parent_token=metadata.retrieve_parent_token, retrieve_parent_token=metadata.retrieve_parent_token,
@@ -621,6 +625,7 @@ class MambaMixer2(torch.nn.Module):
intermediate_states_buffer=layer_cache.intermediate_ssm, intermediate_states_buffer=layer_cache.intermediate_ssm,
cache_steps=draft_token_num, cache_steps=draft_token_num,
retrieve_parent_token=metadata.retrieve_parent_token, retrieve_parent_token=metadata.retrieve_parent_token,
intermediate_state_indices=self.intermediate_state_indices,
) )
else: else:
selective_state_update( selective_state_update(
@@ -56,6 +56,14 @@ else:
is not None 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"]) @triton.jit(do_not_specialize=["T"])
def _selective_scan_update_kernel( def _selective_scan_update_kernel(
# Pointers to matrices # Pointers to matrices
@@ -74,6 +82,7 @@ def _selective_scan_update_kernel(
intermediate_states_buffer, intermediate_states_buffer,
cache_steps, cache_steps,
retrieve_parent_token_ptr, retrieve_parent_token_ptr,
intermediate_state_indices_ptr,
# Matrix dimensions # Matrix dimensions
batch, batch,
T, T,
@@ -130,6 +139,7 @@ def _selective_scan_update_kernel(
DISABLE_STATE_UPDATE: tl.constexpr, DISABLE_STATE_UPDATE: tl.constexpr,
CACHE_INTERMEDIATE_STATES: tl.constexpr, CACHE_INTERMEDIATE_STATES: tl.constexpr,
HAS_EAGLE_TREE_CUSTOM_ATTN_MASK: tl.constexpr, HAS_EAGLE_TREE_CUSTOM_ATTN_MASK: tl.constexpr,
HAS_INTERMEDIATE_STATE_INDICES: tl.constexpr,
BLOCK_SIZE_DSTATE: tl.constexpr, BLOCK_SIZE_DSTATE: tl.constexpr,
): ):
pid_m = tl.program_id(axis=0) pid_m = tl.program_id(axis=0)
@@ -177,7 +187,12 @@ def _selective_scan_update_kernel(
cache_idx = -1 cache_idx = -1
if CACHE_INTERMEDIATE_STATES: 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 cache_idx = state_batch_idx
else: else:
cache_idx = pid_b cache_idx = pid_b
@@ -250,7 +265,7 @@ def _selective_scan_update_kernel(
if state_batch_idx != pad_slot_id: if state_batch_idx != pad_slot_id:
cache_ptr_base = ( cache_ptr_base = (
intermediate_states_buffer intermediate_states_buffer
+ state_batch_idx * cache_steps * nheads * dim * dstate + cache_idx * cache_steps * nheads * dim * dstate
+ current_step_idx * nheads * dim * dstate + current_step_idx * nheads * dim * dstate
+ pid_h * dim * dstate + pid_h * dim * dstate
) )
@@ -300,6 +315,7 @@ def selective_state_update(
intermediate_states_buffer=None, intermediate_states_buffer=None,
cache_steps=None, cache_steps=None,
retrieve_parent_token=None, retrieve_parent_token=None,
intermediate_state_indices=None,
): ):
""" """
Argument: Argument:
@@ -324,6 +340,8 @@ def selective_state_update(
intermediate_states_buffer: Buffer to cache intermediate states intermediate_states_buffer: Buffer to cache intermediate states
cache_steps: Total number of steps in the buffer cache_steps: Total number of steps in the buffer
retrieve_parent_token: (batch, T) tensor of parent token indices for EAGLE tree attention 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: if state.dim() == 3:
state = state.unsqueeze(1) state = state.unsqueeze(1)
@@ -426,6 +444,7 @@ def selective_state_update(
intermediate_states_buffer, intermediate_states_buffer,
cache_steps if cache_steps is not None else 0, cache_steps if cache_steps is not None else 0,
retrieve_parent_token, retrieve_parent_token,
intermediate_state_indices,
batch, batch,
T, T,
nheads, nheads,
+7 -1
View File
@@ -1376,7 +1376,13 @@ class ServerArgs:
else: else:
self.quantization = model_config.quantization self.quantization = model_config.quantization
self.moe_runner_backend = "flashinfer_cutlass" 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( logger.warning(
"Disabling overlap schedule since MambaRadixCache is not compatible with " "Disabling overlap schedule since MambaRadixCache is not compatible with "
"overlap schedule currently, try to use --disable-radix-cache if overlap schedule is necessary" "overlap schedule currently, try to use --disable-radix-cache if overlap schedule is necessary"