[nemetron/mtp] fix nemotron mtp (#16275)
This commit is contained in:
@@ -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,
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
Reference in New Issue
Block a user