diff --git a/examples/inference/gpt/gpt_dynamic_inference.py b/examples/inference/gpt/gpt_dynamic_inference.py index c456f7ea289..251aa100cba 100644 --- a/examples/inference/gpt/gpt_dynamic_inference.py +++ b/examples/inference/gpt/gpt_dynamic_inference.py @@ -11,7 +11,7 @@ from collections import defaultdict from functools import partial from tqdm import tqdm -from typing import Dict, List, Optional +from typing import Dict, List, Tuple, Optional import torch from tqdm import tqdm @@ -28,18 +28,21 @@ from megatron.core.inference.text_generation_controllers.text_generation_controller import ( TextGenerationController, ) +from megatron.core.ssm.mamba_hybrid_layer_allocation import Symbols from megatron.core.tokenizers.text.utils.build_tokenizer import build_tokenizer from megatron.core.transformer.module import MegatronModule +from megatron.core.utils import get_attr_wrapped_model sys.path.append( os.path.abspath(os.path.join(os.path.dirname(__file__), os.path.pardir, os.path.pardir)) ) from megatron.training import get_args, get_model as _get_model, get_tokenizer, initialize_megatron from megatron.training.checkpointing import load_checkpoint - -from megatron.core.utils import configure_nvtx_profiling from model_provider import model_provider from gpt_builders import gpt_builder +from mamba_builders import mamba_builder + +from megatron.core.utils import configure_nvtx_profiling import json @@ -54,7 +57,6 @@ from megatron.training import get_model as _get_model from megatron.training import get_tokenizer, initialize_megatron from megatron.training.checkpointing import load_checkpoint -from pretrain_gpt import model_provider import torch import io @@ -86,9 +88,16 @@ def get_model() -> MegatronModule: args = get_args() + if args.model_provider == "gpt": + model_builder = gpt_builder + elif args.model_provider == "mamba": + model_builder = mamba_builder + else: + raise ValueError(f"Invalid model provider {args.model_provider}") + # Build model. model = _get_model( - partial(model_provider, gpt_builder), + partial(model_provider, model_builder), wrap_with_ddp=False ) @@ -115,7 +124,10 @@ def get_model() -> MegatronModule: def get_inference_context( requests: List[Request], sampling_params: Optional[SamplingParams] = None, - calculate_max_sequence_length_from_requests: bool = True + calculate_max_sequence_length_from_requests: bool = True, + layer_type_list: Optional[List[str]] = None, + mamba_conv_states_shape: Optional[Tuple[int]] = None, + mamba_ssm_states_shape: Optional[Tuple[int]] = None, ): """The inference context manages the KV cache and other inference state.""" @@ -154,6 +166,9 @@ def get_inference_context( max_tokens_override=args.inference_dynamic_batching_max_tokens_override, tensor_model_parallel_size=args.tensor_model_parallel_size, materialize_only_last_token_logits=not args.return_log_probs, + layer_type_list=layer_type_list, + mamba_conv_states_shape=mamba_conv_states_shape, + mamba_ssm_states_shape=mamba_ssm_states_shape, cache_mla_latent=args.multi_latent_attention and args.cache_mla_latents, kv_lora_rank=args.kv_lora_rank if args.multi_latent_attention else None, qk_pos_emb_head_dim=args.qk_pos_emb_head_dim, @@ -364,21 +379,38 @@ def main(): termination_id=args.termination_id if args.termination_id is not None else tokenizer.eod, ) - # Requests, context, conroller. model = get_model() + + # Layer type list for hybrid models + decoder = get_attr_wrapped_model(model, "decoder") + layer_type_list = getattr(decoder, "layer_type_list", None) + if layer_type_list is not None and Symbols.MAMBA in layer_type_list: + (mamba_conv_states_shape, mamba_ssm_states_shape) = decoder.mamba_state_shapes_per_request() + else: + mamba_conv_states_shape = None + mamba_ssm_states_shape = None + + # Requests, context, controller. requests = build_requests(args, tokenizer, sampling_params) - context = get_inference_context(requests, sampling_params) + context = get_inference_context( + requests, + sampling_params, + layer_type_list=layer_type_list, + mamba_conv_states_shape=mamba_conv_states_shape, + mamba_ssm_states_shape=mamba_ssm_states_shape, + ) controller = get_inference_controller(model, context) # Validate all context_length's <= max_tokens. - invalid_prompt_length_map = {} - for request_idx, request in enumerate(requests): - if len(request.prompt_tokens) > context.max_tokens: - invalid_prompt_length_map[request_idx] = len(request.prompt_tokens) - assert not invalid_prompt_length_map, ( - "request idxs with prompts longer than context.max_tokens: " - ", ".join(f"{k}({v})" for k, v in invalid_prompt_length_map.items()) - ) + if args.disable_chunked_prefill: + invalid_prompt_length_map = {} + for request_idx, request in enumerate(requests): + if len(request.prompt_tokens) > context.max_tokens: + invalid_prompt_length_map[request_idx] = len(request.prompt_tokens) + assert not invalid_prompt_length_map, ( + "request idxs with prompts longer than context.max_tokens: " + ", ".join(f"{k}({v})" for k, v in invalid_prompt_length_map.items()) + ) # Inference engine. engine = DynamicInferenceEngine( @@ -418,8 +450,8 @@ def main(): ) # Print unique prompts + outputs. - if torch.distributed.get_rank() == 0: + if torch.distributed.get_rank() == 0: def escape_str(s): return s.replace("\n", "\\n") diff --git a/examples/inference/gpt/utils.py b/examples/inference/gpt/utils.py index baa25787e83..0ea1f5a3df0 100644 --- a/examples/inference/gpt/utils.py +++ b/examples/inference/gpt/utils.py @@ -222,6 +222,9 @@ def arrival(r): if len(time_offsets) == 0: time_offsets = [0.0] + # Ensure first time is 0. + time_offsets = [to - time_offsets[0] for to in time_offsets] + # Truncate to num_requests. assert len(time_offsets) >= num_requests time_offsets = time_offsets[:num_requests] diff --git a/megatron/core/inference/contexts/attention_context/mamba_metadata.py b/megatron/core/inference/contexts/attention_context/mamba_metadata.py new file mode 100644 index 00000000000..e9cd99a6c48 --- /dev/null +++ b/megatron/core/inference/contexts/attention_context/mamba_metadata.py @@ -0,0 +1,106 @@ +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import torch + + +class MambaMetadata: + """Manages the metadata tensors required for Mamba layers during inference.""" + + def __init__(self, max_requests: int): + """ + Initializes the Mamba slot allocator. + + Args: + max_requests (int): The maximum number of concurrent requests. + """ + self.max_requests = max_requests + + # Metadata for mapping requests to slots in the static Mamba state buffer + self.request_to_mamba_state_idx = torch.full( + (self.max_requests,), -1, dtype=torch.int32, device=torch.cuda.current_device() + ) + + # Separate mapping used only for CUDA graph compatibility + self.request_to_mamba_state_idx_cudagraph_only = torch.full( + (self.max_requests,), -1, dtype=torch.int32, device=torch.cuda.current_device() + ) + + # Allocator for Mamba state slots + self.mamba_state_free_slots = torch.arange( + self.max_requests, dtype=torch.int32, device=torch.cuda.current_device() + ) + self.mamba_state_free_slot_count = self.max_requests + + def reset(self) -> None: + """ + Resets all Mamba states and frees all allocated slots. + """ + self.request_to_mamba_state_idx.fill_(-1) + self.request_to_mamba_state_idx_cudagraph_only.fill_(-1) + + # Re-initialize the free slot pool + self.mamba_state_free_slots = torch.arange( + self.max_requests, dtype=torch.int32, device=torch.cuda.current_device() + ) + self.mamba_state_free_slot_count = self.max_requests + + def reset_cudagraph_mapping(self) -> None: + """ + Resets only the CUDA graph mapping tensor. + """ + self.request_to_mamba_state_idx_cudagraph_only.fill_(-1) + + def update_cudagraph_mapping( + self, active_mamba_indices: torch.Tensor, num_active_requests: int + ) -> None: + """ + Updates the dedicated CUDA graph mapping tensor with the indices + of currently active requests. + + Args: + active_mamba_indices (Tensor): Tensor containing the Mamba slot indices + for active requests. + num_active_requests (int): The number of active requests. + """ + self.request_to_mamba_state_idx_cudagraph_only[0:num_active_requests] = active_mamba_indices + + def allocate_slot(self) -> int: + """ + Allocates a new slot for a request in the Mamba state buffers. + + Returns: + int: The index of the allocated slot. + Returns None if no slots are available. + """ + if self.mamba_state_free_slot_count == 0: + return None + + # Get a free slot + self.mamba_state_free_slot_count -= 1 + mamba_idx = self.mamba_state_free_slots[self.mamba_state_free_slot_count] + + return mamba_idx + + def free_slots(self, request_indices: torch.Tensor) -> None: + """ + Frees the Mamba state slots associated with the given request indices. + + Args: + request_indices (Tensor): A 1D tensor of request indices to free. + """ + # Get the Mamba state indices for finished requests + mamba_indices_to_free = self.request_to_mamba_state_idx[request_indices] + + # Filter out any invalid indices (e.g., -1) + mamba_indices_to_free = mamba_indices_to_free[mamba_indices_to_free != -1] + num_to_free = len(mamba_indices_to_free) + + if num_to_free > 0: + # Add the freed indices back to the free slot pool + start_idx = self.mamba_state_free_slot_count + end_idx = start_idx + num_to_free + self.mamba_state_free_slots[start_idx:end_idx] = mamba_indices_to_free + self.mamba_state_free_slot_count = end_idx + + # Invalidate the Mamba state index for the finished requests + self.request_to_mamba_state_idx[request_indices] = -1 diff --git a/megatron/core/inference/contexts/dynamic_context.py b/megatron/core/inference/contexts/dynamic_context.py index d6cc2598998..000b58200f8 100644 --- a/megatron/core/inference/contexts/dynamic_context.py +++ b/megatron/core/inference/contexts/dynamic_context.py @@ -23,9 +23,14 @@ from megatron.core.inference.utils import tensor_swap from megatron.core.models.common.embeddings.rope_utils import apply_rotary_pos_emb from megatron.core.package_info import __version__ as mcore_version +from megatron.core.ssm.mamba_hybrid_layer_allocation import ( + Symbols, + get_layer_maps_from_layer_type_list, +) from megatron.core.transformer import TransformerConfig from megatron.core.utils import divide as core_divide +from .attention_context.mamba_metadata import MambaMetadata from .attention_context.mha_metadata import GraphedMHAMetadata, NonGraphedMHAMetadata from .base_context import BaseInferenceContext from .dynamic_block_allocator import BlockAllocator @@ -227,8 +232,17 @@ class DynamicInferenceContext(BaseInferenceContext): where the cuda graph batch sizes range from 1 to `max_requests` (as computed below). Due to rounding, the actual number of cuda graphs may not equal this argument. - materialize_only_last_token_logits (bool): If True, only the last token logits - are materialized in the context. + materialize_only_last_token_logits (Optional[bool]): Whether to only + materialize logits for the last token. This should be set to False + if returning log probs. + layer_type_list (Optional[List[str]]): A list of strings that indicates + the layer type (Mamba / Attention / MLP) for each layer. + See `megatron/core/ssm/mamba_hybrid_layer_allocation.py` for the list + of symbols. This must be provided for hybrid models. + mamba_conv_states_shape: (Optional[Tuple[int]]): Mamba conv states shape per request. + This must be provided for hybrid models. + mamba_ssm_states_shape: (Optional[Tuple[int]]): Mamba ssm states shape per request. + This must be provided for hybrid models. use_cuda_graphs_for_non_decode_steps (bool): If True, use cuda graphs for non-decode engine steps. unified_memory_level (Optional[int]): Set unified memory usage within the @@ -259,7 +273,10 @@ def __init__( kv_lora_rank: Optional[int] = None, qk_pos_emb_head_dim: Optional[int] = None, num_cuda_graphs: Optional[int] = None, - materialize_only_last_token_logits: bool = True, + materialize_only_last_token_logits: Optional[bool] = True, + layer_type_list: Optional[List[str]] = None, + mamba_conv_states_shape: Optional[Tuple[int]] = None, + mamba_ssm_states_shape: Optional[Tuple[int]] = None, use_cuda_graphs_for_non_decode_steps: bool = True, use_flashinfer_fused_rope: bool = False, unified_memory_level: Optional[int] = 0, @@ -283,6 +300,41 @@ def __init__( tp_size = tensor_model_parallel_size hidden_size_per_attention_head = core_divide(projection_size, num_attention_heads) num_attention_heads_per_partition = core_divide(num_attention_heads, tp_size) + + # Mamba states. + self.is_hybrid_model = layer_type_list is not None and Symbols.MAMBA in layer_type_list + if self.is_hybrid_model: + assert ( + mamba_conv_states_shape is not None + ), "`mamba_conv_states_shape` must be specified for hybrid models" + assert ( + mamba_ssm_states_shape is not None + ), "`mamba_ssm_states_shape` must be specified for hybrid models" + assert ( + not use_cuda_graphs_for_non_decode_steps + ), "Non-decode CUDA graphs not yet supported for hybrid models" + + # For hybrid models, the layer map converts the global layer index to the + # corresponding attention layer index or Mamba layer index depending on the + # layer type. + attention_layer_map, mamba_layer_map, _ = get_layer_maps_from_layer_type_list( + layer_type_list + ) + self.num_attention_layers = len(attention_layer_map) + self.num_mamba_layers = len(mamba_layer_map) + self.layer_map = attention_layer_map | mamba_layer_map + else: + # The layer map is the identity function for pure Transformer models. + self.num_attention_layers = num_layers + self.num_mamba_layers = 0 + (mamba_conv_states_shape, mamba_ssm_states_shape) = (None, None) + self.layer_map = {i: i for i in range(self.num_attention_layers)} + + if self.num_attention_layers == 0: + raise NotImplementedError( + f"Using `DynamicInferenceContext` with no attention is not supported." + ) + # Block size tokens, bytes. dtype_size_bytes = params_dtype.itemsize self.block_size_tokens = block_size_tokens @@ -297,24 +349,38 @@ def __init__( self.block_size_bytes = ( dtype_size_bytes * 2 # key, value - * num_layers + * self.num_attention_layers * self.block_size_tokens * num_attention_heads_per_partition * hidden_size_per_attention_head ) + assert self.block_size_bytes > 0 # Adjust buffer to be a multiple of block size. buffer_size_bytes = int(buffer_size_gb * 1024**3) buffer_size_bytes_rem = buffer_size_bytes % self.block_size_bytes buffer_size_bytes = buffer_size_bytes - buffer_size_bytes_rem - # Compute max_requets, max_tokens from buffer size and overflow factor. + mamba_states_memory_per_request = 0 + if self.is_hybrid_model: + mamba_states_memory_per_request += math.prod(mamba_conv_states_shape) + mamba_states_memory_per_request += math.prod(mamba_ssm_states_shape) + mamba_states_memory_per_request *= self.num_mamba_layers + mamba_states_memory_per_request *= dtype_size_bytes + + # Compute max_requets, max_tokens from buffer size, overflow factor, and Mamba state size. def bytes_to_max_requests_and_tokens(n_bytes): - n_tokens = n_bytes / self.block_size_bytes * self.block_size_tokens - n_requests = n_tokens / max_sequence_length - return self.round_up_requests(int(n_requests), tp_size=tp_size), self.round_up_tokens( - int(n_tokens), tp_size=tp_size + bytes_per_token = self.block_size_bytes / self.block_size_tokens + cost_per_request_bytes = ( + mamba_states_memory_per_request + max_sequence_length * bytes_per_token ) + # TODO(ksanthanam): Leave room for an extra request in the event of padding + # for non-decode CUDA graphs + n_requests = n_bytes / cost_per_request_bytes + n_tokens = n_requests * max_sequence_length + n_requests = self.round_up_requests(int(n_requests), tp_size=tp_size) + n_tokens = self.round_up_tokens(int(n_tokens), tp_size=tp_size) + return n_requests, n_tokens self.max_requests, self.max_tokens = bytes_to_max_requests_and_tokens(buffer_size_bytes) if buffer_overflow_factor is not None: @@ -339,7 +405,6 @@ def bytes_to_max_requests_and_tokens(n_bytes): # Initialize context state. self.params_dtype = params_dtype - self.num_layers = num_layers self.max_sequence_length = max_sequence_length # Unified memory. @@ -390,8 +455,11 @@ def bytes_to_max_requests_and_tokens(n_bytes): self.token_to_position_in_request = torch.empty_like(self.token_to_input_ids) self.token_to_local_position_within_kv_block = torch.empty_like(self.token_to_input_ids) - # Calculate the total number of blocks available in the buffer - block_count_total = buffer_size_bytes // self.block_size_bytes + # Calculate the total number of chunks available in the buffer + total_mamba_states_memory = mamba_states_memory_per_request * self.max_requests + block_count_total = ( + max(0, buffer_size_bytes - total_mamba_states_memory) // self.block_size_bytes + ) # Memory buffer. ctx_manager = ( @@ -402,7 +470,12 @@ def bytes_to_max_requests_and_tokens(n_bytes): with ctx_manager: if cache_mla_latent: self.memory_buffer = torch.full( - (self.num_layers, block_count_total, self.block_size_tokens, kv_reduced_dim), + ( + self.num_attention_layers, + block_count_total, + self.block_size_tokens, + kv_reduced_dim, + ), -1, dtype=self.params_dtype, device=torch.cuda.current_device(), @@ -411,7 +484,7 @@ def bytes_to_max_requests_and_tokens(n_bytes): self.memory_buffer = torch.full( ( 2, # key and value - self.num_layers, + self.num_attention_layers, block_count_total, self.block_size_tokens, num_attention_heads_per_partition, @@ -516,14 +589,34 @@ def bytes_to_max_requests_and_tokens(n_bytes): block_count_total=block_count_total, gtd_block_count=self.gtd_block_count ) + # Optional state tensors for hybrid models + if self.is_hybrid_model: + self.mamba_metadata = MambaMetadata(max_requests=self.max_requests) + + with ctx_manager: + self.mamba_conv_states = torch.zeros( + (self.num_mamba_layers, self.max_requests) + mamba_conv_states_shape, + dtype=self.params_dtype, + device=torch.cuda.current_device(), + ) + self.mamba_ssm_states = torch.zeros( + (self.num_mamba_layers, self.max_requests) + mamba_ssm_states_shape, + dtype=self.params_dtype, + device=torch.cuda.current_device(), + ) + + else: + self.mamba_metadata = None + # Store the dummy block idx reference for convenience self.dummy_block_idx = self.block_allocator.dummy_block_idx # Deal with chunked prefill self.chunked_prefill_request_id = -1 - # Reset attention state. + # Reset attention and Mamba state. self.reset_attention_state() + self.reset_mamba_state() if use_flashinfer_fused_rope is True: assert HAVE_FLASHINFER, "flashinfer is not installed" @@ -628,7 +721,8 @@ def is_decode_only(self) -> bool: """Test if all active requests are in decode phase. For a request in prefill phase active_tokens = query length - Once the request moves to decode phase active tokens is 1 for that request. So if all active requests are in decode phase, they will be equal to active token count. + Once the request moves to decode phase active tokens is 1 for that request. + So if all active requests are in decode phase, they will be equal to active token count. """ total_active_requests = self.total_request_count - self.paused_request_count return total_active_requests == self.active_token_count @@ -664,11 +758,7 @@ def get_max_sequence_lengths(self) -> Tensor: def get_active_request_count(self): """Returns the current number of active requests.""" - active_sequence_lengths = self.get_active_sequence_lengths() - max_sequence_lengths = self.get_max_sequence_lengths() - active_requests_mask = torch.less(active_sequence_lengths, max_sequence_lengths).byte() - active_request_count = (active_requests_mask == 1).sum().item() - return active_request_count + return self.total_request_count - self.paused_request_count def append_key_value_cache(self, layer_number: int, key: Tensor, value: Tensor) -> None: """Append to KV cache. @@ -678,10 +768,12 @@ def append_key_value_cache(self, layer_number: int, key: Tensor, value: Tensor) key (Tensor): Key tensor. value (Tensor): Value tensor. """ + attention_layer_number = self.layer_map[layer_number - 1] + if triton_append_key_value_cache is not None and not self.cache_mla_latent: # currently does not support MLA latent cache return triton_append_key_value_cache( - layer_number=layer_number, + layer_number=attention_layer_number, key=key, value=value, memory_buffer=self.memory_buffer, @@ -706,14 +798,14 @@ def append_key_value_cache(self, layer_number: int, key: Tensor, value: Tensor) if self.cache_mla_latent: # We pass the kv_concat as the key in cache_mla_latent kv_concat = key - self.memory_buffer[layer_number - 1, block_idx, local_kv_seq_idx] = kv_concat[ + self.memory_buffer[attention_layer_number, block_idx, local_kv_seq_idx] = kv_concat[ : self.padded_active_token_count ] else: - self.memory_buffer[0, layer_number - 1, block_idx, local_kv_seq_idx] = key[ + self.memory_buffer[0, attention_layer_number, block_idx, local_kv_seq_idx] = key[ : self.padded_active_token_count ] - self.memory_buffer[1, layer_number - 1, block_idx, local_kv_seq_idx] = value[ + self.memory_buffer[1, attention_layer_number, block_idx, local_kv_seq_idx] = value[ : self.padded_active_token_count ] @@ -727,19 +819,30 @@ def key_value_cache(self, layer_number: int) -> Tuple[Tensor, Tensor]: (Tuple[Tensor, Tensor]) The key and value pointer tensors that point to blocks within the block-level memory buffer. """ + attention_layer_number = self.layer_map[layer_number - 1] if self.cache_mla_latent: return ( - self.memory_buffer[layer_number - 1], + self.memory_buffer[attention_layer_number], None, self.active_attn_metadata["mha_metadata"].state_data["block_table"], ) else: return ( - self.memory_buffer[0, layer_number - 1], - self.memory_buffer[1, layer_number - 1], + self.memory_buffer[0, attention_layer_number], + self.memory_buffer[1, attention_layer_number], self.active_attn_metadata["mha_metadata"].state_data["block_table"], ) + def mamba_states_cache(self, layer_number: int) -> Tuple[Tensor, Tensor]: + """Returns the Mamba state tensors for the given layer.""" + assert self.is_hybrid_model, "Only hybrid models have Mamba state tensors" + + mamba_layer_number = self.layer_map[layer_number - 1] + conv_state = self.mamba_conv_states[mamba_layer_number] + ssm_state = self.mamba_ssm_states[mamba_layer_number] + + return (conv_state, ssm_state) + def apply_fused_qk_rotary_emb( self, query: Tensor, key: Tensor, cos_sin_emb: Tensor, config: TransformerConfig ) -> Tuple[Tensor, Tensor]: @@ -854,6 +957,16 @@ def reset_attention_state(self) -> None: attn_metadata.reset() self.active_attn_metadata = None + if self.is_hybrid_model: + self.mamba_metadata.reset_cudagraph_mapping() + + def reset_mamba_state(self) -> None: + """Reset state used within Mamba layers.""" + if self.is_hybrid_model: + self.mamba_conv_states.fill_(0) + self.mamba_ssm_states.fill_(0) + self.mamba_metadata.reset() + def using_cuda_graph_this_step(self) -> bool: """Returns True if cuda graphs are being used for this step.""" has_cuda_graphs = self.cuda_graph_token_counts is not None @@ -977,6 +1090,17 @@ def initialize_attention_state( ) # All attention metadata calculations are now handled by MHAMetadata.update() + # Create Mamba state block table if it's a hybrid model + if self.is_hybrid_model: + active_mamba_indices = self.mamba_metadata.request_to_mamba_state_idx[ + self.paused_request_count : self.total_request_count + ] + + if self.is_decode_only() or self.using_cuda_graph_this_step(): + self.mamba_metadata.update_cudagraph_mapping( + active_mamba_indices, self.total_request_count - self.paused_request_count + ) + def reset(self) -> None: """Reset entire context. @@ -1018,15 +1142,13 @@ def reset(self) -> None: # Reset available block count. self.reset_attention_state() + self.reset_mamba_state() self.block_allocator.reset() self.request_to_kv_block_ids.fill_(-1) # Reset chunked prefill state self.chunked_prefill_request_id = -1 - # Reset chunked prefill state - self.chunked_prefill_request_id = -1 - def current_input_and_position_ids( self, *, num_warmup_tokens: Optional[int] = None ) -> Tuple[Tensor, Tensor]: @@ -1198,6 +1320,18 @@ def add_request(self, req: DynamicInferenceRequest, chunk_length: Optional[int] self.token_to_local_position_within_kv_block[ self.active_token_count : self.active_token_count + chunk_length ] = (token_offset_range % self.block_size_tokens) + + if self.is_hybrid_model and not is_chunked_prefill: + # Allocate a slot for Mamba states + mamba_idx = self.mamba_metadata.allocate_slot() + if mamba_idx is None: + raise ContextOverflowError(req.request_id, "No Mamba slots available") + + # Initialize the allocated Mamba state + self.mamba_conv_states[:, mamba_idx] = 0.0 + self.mamba_ssm_states[:, mamba_idx] = 0.0 + self.mamba_metadata.request_to_mamba_state_idx[self.total_request_count] = mamba_idx + self.active_token_count += chunk_length self.total_request_count += 0 if req.finished_chunk_token_count > 0 else 1 @@ -1216,6 +1350,11 @@ def _move_book_keeping_tensors(self, src_idxs, dst_idxs, next_tokens): self.request_last_kv_block_id[dst_idxs] = self.request_last_kv_block_id[src_idxs] self.request_last_kv_block_offset[dst_idxs] = self.request_last_kv_block_offset[src_idxs] + if self.is_hybrid_model: + self.mamba_metadata.request_to_mamba_state_idx[dst_idxs] = ( + self.mamba_metadata.request_to_mamba_state_idx[src_idxs] + ) + def _swap_book_keeping_tensors(self, src_idxs, dst_idxs, next_tokens): """ Swaps all the relevent booking tensors with src idxs to dst idxs @@ -1230,6 +1369,9 @@ def _swap_book_keeping_tensors(self, src_idxs, dst_idxs, next_tokens): tensor_swap(self.request_last_kv_block_id, src_idxs, dst_idxs) tensor_swap(self.request_last_kv_block_offset, src_idxs, dst_idxs) + if self.is_hybrid_model: + tensor_swap(self.mamba_metadata.request_to_mamba_state_idx, src_idxs, dst_idxs) + # TODO: see if we can compile this function def update_requests(self, active_requests_mask: Tensor, new_tokens: Tensor) -> Tensor: """Update context state after calling engine.step(). @@ -1301,10 +1443,17 @@ def update_requests(self, active_requests_mask: Tensor, new_tokens: Tensor) -> T non_zero_values_in_kv_memory = kv_blocks_assigned[kv_blocks_assigned != -1] self.block_allocator.release_memory_blocks(non_zero_values_in_kv_memory) + if self.is_hybrid_model: + self.mamba_metadata.free_slots(finished_idxs) + # Reset request/token counts. self.request_to_kv_block_ids.fill_(-1) self.total_request_count = 0 self.active_token_count = 0 + + # Reset Mamba state. + self.reset_mamba_state() + return # 3. Concatenate the paused tokens to the active tokens if present. @@ -1332,6 +1481,10 @@ def update_requests(self, active_requests_mask: Tensor, new_tokens: Tensor) -> T # and updates it instead of the original tensor. self.request_to_kv_block_ids[finished_idxs] = -1 + if self.is_hybrid_model: + # Get the Mamba state indices for finished requests and free them + self.mamba_metadata.free_slots(finished_idxs) + if active_request_count > 0: finished_idxs_on_left = ( torch.nonzero(active_requests_mask[:active_request_count] == 0, as_tuple=True)[ @@ -1351,8 +1504,10 @@ def update_requests(self, active_requests_mask: Tensor, new_tokens: Tensor) -> T next_tokens=next_tokens, ) - # Reset block ids for recently moved requests. + # Reset chunk ids for recently moved requests. self.request_to_kv_block_ids[active_idxs_on_right] = -1 + if self.is_hybrid_model: + self.mamba_metadata.request_to_mamba_state_idx[active_idxs_on_right] = -1 # 5. We identify requests that require a new block and add them to the paused requests (i.e move them left) :- # a) Put requests that have filled their current block and require a new one in a pause state temporarily @@ -1450,6 +1605,7 @@ def update_requests(self, active_requests_mask: Tensor, new_tokens: Tensor) -> T # 7. We make changes to the request book keeping tesnsors and setup the tokens for next iteration self.total_request_count = active_request_count + self.paused_request_count + # All these active requests are in decode phase, so they need only 1 token per request self.active_token_count = active_request_count # Always the first section of token input ids are only used. diff --git a/megatron/core/inference/contexts/fused_kv_append_kernel.py b/megatron/core/inference/contexts/fused_kv_append_kernel.py index 2078878c8f4..db1eed456e1 100644 --- a/megatron/core/inference/contexts/fused_kv_append_kernel.py +++ b/megatron/core/inference/contexts/fused_kv_append_kernel.py @@ -119,8 +119,8 @@ def triton_append_key_value_cache( _, num_heads, h_dim = key.shape - key_cache = memory_buffer[0, layer_number - 1] - value_cache = memory_buffer[1, layer_number - 1] + key_cache = memory_buffer[0, layer_number] + value_cache = memory_buffer[1, layer_number] key_to_cache = key[:n_tokens] value_to_cache = value[:n_tokens] diff --git a/megatron/core/inference/engines/dynamic_engine.py b/megatron/core/inference/engines/dynamic_engine.py index bcde4f9894d..2c43a7e2611 100644 --- a/megatron/core/inference/engines/dynamic_engine.py +++ b/megatron/core/inference/engines/dynamic_engine.py @@ -702,7 +702,7 @@ def schedule_chunked_prefill(self): # is_continuing_chunked_prefill is True if we are scheduling next # chunk of a existing chunked prefill request - is_continuing_chunked_prefill = self.context.chunked_prefill_request_id > 0 + is_continuing_chunked_prefill = self.context.chunked_prefill_request_id >= 0 # Use remaining prompt tokens for scheduling decisions remaining_len = len(req.remaining_prompt_tokens) @@ -939,7 +939,7 @@ def generate( result = self.step_modern() finished_requests_list.extend(result["finished_requests"]) - # Ensure requests are returned in the same order they were passed in. + # Ensure requests are returned in the same order they were passed in finished_requests_list.sort(key=lambda x: x.request_id) return finished_requests_list diff --git a/megatron/core/models/gpt/gpt_model.py b/megatron/core/models/gpt/gpt_model.py index dbc5a88fc81..25546d36629 100644 --- a/megatron/core/models/gpt/gpt_model.py +++ b/megatron/core/models/gpt/gpt_model.py @@ -588,8 +588,6 @@ def _postprocess( # Perform the sequence parallel gather here instead of after the output layer # because we need to slice the last token logits from the full view of the # packed logits across all requests. - # TODO(ksanthanam): Make the equivalent change in the `MambaModel` code after - # merging in !3722. hidden_states = gather_from_sequence_parallel_region( hidden_states, group=self.pg_collection.tp ) @@ -597,7 +595,7 @@ def _postprocess( sequence_parallel_override = True # Reshape [B, 1, H] to [1, B, H] → extract each sample’s true last‐token hidden - # state ([B, H]) → unsqueeze back to [1, B, H] + # state ([B, H]) → unsqueeze back to [B, 1, H] # (so that the output layer, which expects S×B×H, receives only the final token) hidden_states = inference_context.last_token_logits( hidden_states.squeeze(1).unsqueeze(0) diff --git a/megatron/core/models/mamba/mamba_model.py b/megatron/core/models/mamba/mamba_model.py index fb3df5e23f2..378cf7e47d6 100644 --- a/megatron/core/models/mamba/mamba_model.py +++ b/megatron/core/models/mamba/mamba_model.py @@ -12,6 +12,7 @@ from megatron.core.models.common.language_module.language_module import LanguageModule from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.quantization.utils import get_quant_config_or_none +from megatron.core.tensor_parallel import gather_from_sequence_parallel_region from megatron.core.transformer import TransformerConfig from megatron.core.transformer.enums import ModelType from megatron.core.transformer.spec_utils import ModuleSpec, build_module @@ -244,13 +245,41 @@ def forward( if self.share_embeddings_and_output_weights: output_weight = self.shared_embedding_or_output_weight() + sequence_parallel_override = False if in_inference_mode and inference_context.materialize_only_last_token_logits: - hidden_states = hidden_states[-1, :, :].unsqueeze(0) + if inference_context.is_static_batching(): + hidden_states = hidden_states[-1:, :, :] + else: + if self.output_layer.sequence_parallel: + # Perform the sequence parallel gather here instead of after the output layer + # because we need to slice the last token logits from the full view of the + # packed logits across all requests. + hidden_states = gather_from_sequence_parallel_region( + hidden_states, group=self.pg_collection.tp + ) + self.output_layer.sequence_parallel = False + sequence_parallel_override = True + + # Reshape [B, 1, H] to [1, B, H] → extract each sample’s true last‐token hidden + # state ([B, H]) → unsqueeze back to [B, 1, H] + # (so that the output layer, which expects S×B×H, receives only the final token) + hidden_states = inference_context.last_token_logits( + hidden_states.squeeze(1).unsqueeze(0) + ).unsqueeze(1) logits, _ = self.output_layer( hidden_states, weight=output_weight, runtime_gather_output=runtime_gather_output ) + # Restore sequence parallel execution to the output layer if necessary. + if sequence_parallel_override: + assert ( + in_inference_mode + and inference_context.is_dynamic_batching() + and inference_context.materialize_only_last_token_logits + ) + self.output_layer.sequence_parallel = True + if labels is None: # [s b h] => [b s h] return logits.transpose(0, 1).contiguous() diff --git a/megatron/core/ssm/mamba_block.py b/megatron/core/ssm/mamba_block.py index 01b9f4eac66..7d8ca74c8f2 100644 --- a/megatron/core/ssm/mamba_block.py +++ b/megatron/core/ssm/mamba_block.py @@ -9,7 +9,7 @@ from contextlib import nullcontext from dataclasses import dataclass from functools import partial -from typing import Optional, Union +from typing import Optional, Tuple, Union import torch from torch import Tensor, nn @@ -147,7 +147,7 @@ def __init__( self.hybrid_mlp_ratio = hybrid_mlp_ratio self.hybrid_override_pattern = hybrid_override_pattern - layer_type_list = allocate_layers( + self.layer_type_list = allocate_layers( self.config.num_layers, self.hybrid_attention_ratio, self.hybrid_mlp_ratio, @@ -156,12 +156,12 @@ def __init__( pp_layer_offset = 0 if self.pp_group.size() > 1: - pp_layer_offset, layer_type_list = self._select_layers_for_pipeline_parallel( - layer_type_list + pp_layer_offset, self.layer_type_list = self._select_layers_for_pipeline_parallel( + self.layer_type_list ) self.layers = nn.ModuleList() - for i, layer_type in enumerate(layer_type_list): + for i, layer_type in enumerate(self.layer_type_list): fp8_init_context = get_fp8_context(self.config, i + pp_layer_offset, is_init=True) with fp8_init_context: if layer_type == LayerSymbols.MAMBA: @@ -224,22 +224,6 @@ def _select_layers_for_pipeline_parallel(self, layer_type_list): return offset, selected_list - def allocate_inference_cache(self, batch_size, max_seqlen, dtype=None): - """ - Allocate inference cache for each layer. - - Args: - batch_size (int): The batch size to use for inference. - max_seqlen (int): The maximum sequence length to use - for inference. - dtype (optional): The data type to use for allocation. - Defaults to the data type of the model. - """ - return { - i: layer.allocate_inference_cache(batch_size, max_seqlen, dtype=dtype) - for i, layer in enumerate(self.layers) - } - def set_input_tensor(self, input_tensor: Tensor): """Set input tensor to be used instead of forward()'s input. @@ -250,6 +234,16 @@ def set_input_tensor(self, input_tensor: Tensor): forward_step_func""" self.input_tensor = input_tensor + def mamba_state_shapes_per_request(self) -> Optional[Tuple[Tuple[int], Tuple[int]]]: + """ + Returns the Mamba conv and ssm states shapes per input sequence + if this block contains Mamba layers (this may not be the case with PP > 1). + """ + for layer_type, layer in zip(self.layer_type_list, self.layers): + if layer_type == LayerSymbols.MAMBA: + return layer.mamba_state_shapes_per_request() + return None + def forward( self, hidden_states: Union[Tensor, WrappedTensor], @@ -287,10 +281,7 @@ def forward( if isinstance(hidden_states, WrappedTensor): hidden_states = hidden_states.unwrap() - if inference_context: - assert ( - inference_context.is_static_batching() - ), "Mamba currently does not support dynamic inference batching." + if inference_context and inference_context.is_static_batching(): # NOTE(bnorick): match BaseInferenceContext attributes for # mamba_ssm.utils.generation.BaseInferenceContext, # this hack supports eval diff --git a/megatron/core/ssm/mamba_hybrid_layer_allocation.py b/megatron/core/ssm/mamba_hybrid_layer_allocation.py index 26972b5454b..7407bfe899f 100644 --- a/megatron/core/ssm/mamba_hybrid_layer_allocation.py +++ b/megatron/core/ssm/mamba_hybrid_layer_allocation.py @@ -1,20 +1,30 @@ # Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. import logging +from typing import Dict, List, Tuple if __name__ != "__main__": from megatron.core.utils import log_single_rank else: from typing import Any + import torch + def log_single_rank(logger: logging.Logger, *args: Any, rank: int = 0, **kwargs: Any): - print(*args[1:], **kwargs) + """Logs a message to the given rank.""" + if torch.distributed.is_initialized(): + if torch.distributed.get_rank() == rank: + logger.log(*args, **kwargs) + else: + logger.log(*args, **kwargs) logger = logging.getLogger(__name__) class Symbols: + """Symbols for different layer types.""" + MAMBA = "M" ATTENTION = "*" MLP = "-" @@ -87,6 +97,7 @@ def allocate_layers( target_mlp_ratio: float, override_pattern: str = None, ) -> list: + """Allocates layers according to the requested distribution of layer types.""" assert total_layers_count > 0 assert target_attention_ratio >= 0.0 and target_attention_ratio <= 1.0 assert target_mlp_ratio >= 0.0 and target_mlp_ratio <= 1.0 @@ -156,6 +167,22 @@ def allocate_layers( return layer_type_list +def get_layer_maps_from_layer_type_list( + layer_type_list: List[str], +) -> Tuple[Dict[int, int], Dict[int, int], Dict[int, int]]: + """ + Returns maps from global layer index to the corresponding layer index + for each layer type in [Attention, Mamba, MLP] given a layer type list. + """ + layer_types = [Symbols.ATTENTION, Symbols.MAMBA, Symbols.MLP] + layer_maps = {layer_type: {} for layer_type in layer_types} + for global_layer_idx, layer_type in enumerate(layer_type_list): + layer_map = layer_maps[layer_type] + local_layer_idx = len(layer_map) + layer_map[global_layer_idx] = local_layer_idx + return [layer_maps[layer_type] for layer_type in layer_types] + + if __name__ == "__main__": test_cases = [ # (10, 0.2, 0.0), @@ -187,5 +214,5 @@ def allocate_layers( (9, 0.0, 0.0, "MMMMMMMMM"), ] for t in test_cases: - print("") + logging.info("") allocate_layers(*t) diff --git a/megatron/core/ssm/mamba_layer.py b/megatron/core/ssm/mamba_layer.py index d83d518331c..69d5ef21c81 100644 --- a/megatron/core/ssm/mamba_layer.py +++ b/megatron/core/ssm/mamba_layer.py @@ -6,7 +6,7 @@ # LICENSE file in the root directory of this source tree. from dataclasses import dataclass, field -from typing import Dict, Optional, Union +from typing import Dict, Optional, Tuple, Union import torch from torch import Tensor @@ -82,6 +82,10 @@ def __init__( self.mamba_bda = build_module(submodules.mamba_bda) self.bias_dropout_add_exec_handler = torch.enable_grad + def mamba_state_shapes_per_request(self) -> Tuple[Tuple[int], Tuple[int]]: + """Returns the Mamba conv and ssm states shapes per request.""" + return self.mixer.mamba_state_shapes_per_request() + def forward( self, hidden_states: Tensor, @@ -127,10 +131,6 @@ def forward( return hidden_states - def allocate_inference_cache(self, batch_size, max_seqlen, dtype=None): - """Allocate the inference cache.""" - return self.mixer.allocate_inference_cache(batch_size, max_seqlen, dtype=dtype) - def sharded_state_dict( self, prefix: str = '', sharded_offsets: tuple = (), metadata: Optional[dict] = None ) -> ShardedStateDict: diff --git a/megatron/core/ssm/mamba_mixer.py b/megatron/core/ssm/mamba_mixer.py index 2caa36fb1e9..895792ff05e 100644 --- a/megatron/core/ssm/mamba_mixer.py +++ b/megatron/core/ssm/mamba_mixer.py @@ -27,7 +27,12 @@ make_sharded_tensors_for_checkpoint, sharded_state_dict_default, ) -from megatron.core.utils import deprecate_inference_params, log_single_rank +from megatron.core.utils import ( + check_mamba_sequence_packing_support, + deprecate_inference_params, + log_single_rank, + maybe_cat, +) from .mamba_context_parallel import MambaContextParallel @@ -38,6 +43,7 @@ try: from causal_conv1d import causal_conv1d_fn, causal_conv1d_update + from causal_conv1d.causal_conv1d_varlen import causal_conv1d_varlen_states except ImportError: causal_conv1d_fn = None causal_conv1d_update = None @@ -63,7 +69,6 @@ except ImportError: HAVE_EINOPS = False - logger = logging.getLogger(__name__) @@ -392,11 +397,11 @@ def forward( inference_context = deprecate_inference_params(inference_context, inference_params) - _, batch, dim = hidden_states.shape + in_inference_mode = inference_context is not None and not self.training + _, batch, dim = hidden_states.shape conv_state, ssm_state = None, None - in_inference_mode = inference_context is not None and not self.training if in_inference_mode: if inference_context.is_dynamic_batching(): return self.dynamic_inference(hidden_states, inference_context) @@ -424,27 +429,151 @@ def forward( return out, out_bias - def dynamic_inference( - self, hidden_states, inference_context: DynamicInferenceContext - ) -> Tuple[torch.Tensor, torch.Tensor]: - """Runs inference computation for dynamic batching.""" - raise NotImplementedError(f"Dynamic inference is not supported.") + def dynamic_inference(self, hidden_states: torch.Tensor, context: DynamicInferenceContext): + """ + Executes dynamic inference by separating decode and prefill requests and + running them independently. Also runs the chunked prefill request independently + if it exists. + """ + sequence_packing_available, reason_for_no_sequence_packing = ( + check_mamba_sequence_packing_support() + ) + assert sequence_packing_available, reason_for_no_sequence_packing + + conv_state, ssm_state = context.mamba_states_cache(self.layer_number) + + # Fast path: decode-only + if context.is_decode_only(): + batch_indices = context.mamba_metadata.request_to_mamba_state_idx_cudagraph_only[ + : context.padded_active_token_count + ] + out, out_bias = self.decode( + hidden_states, conv_state, ssm_state, batch_indices=batch_indices + ) + return out, out_bias + + # Compute input projection before splitting into prefill and decode + # to ensure sequence parallel all-gather. + zxBCdt, _ = self.in_proj(hidden_states) + + # Compute split between decode and prefill. + seq_idx, cu_seqlens, return_varlen_states = self._get_varlen_generation_state(context) + active_query_lengths = context.request_query_lengths[ + context.paused_request_count : context.total_request_count + ] + batch_indices = context.mamba_metadata.request_to_mamba_state_idx + + # First request with query len > 1 is prefill-start. + first_prefill_token_idx = torch.nonzero(active_query_lengths > 1)[0].int() + + # Process decode requests if there are any. + if first_prefill_token_idx > 0: + zxBCdt_decode = zxBCdt[:first_prefill_token_idx] + batch_indices_decode = batch_indices[:first_prefill_token_idx] + y_decode = self.ssm_decode( + zxBCdt_decode.transpose(0, 1), conv_state, ssm_state, batch_indices_decode + ).transpose(0, 1) + else: + y_decode = None + + active_token_count = context.active_token_count + active_request_count = context.get_active_request_count() + padded_active_token_count = context.padded_active_token_count + + # Process the chunked prefill request if it exists. + if context.chunked_prefill_request_id != -1: + chunked_prefill_request_token_count = active_query_lengths[-1] + zxBCdt_chunked_prefill = zxBCdt[ + active_token_count - chunked_prefill_request_token_count : active_token_count + ] + batch_index_chunked_prefill = batch_indices[context.chunked_prefill_request_id] + + y_prefill_chunked = self.ssm_prefill( + zxBCdt_chunked_prefill, + conv_state=conv_state[batch_index_chunked_prefill].unsqueeze(0), + ssm_state=ssm_state[batch_index_chunked_prefill].unsqueeze(0), + is_chunked_prefill=True, + ) - def decode(self, hidden_states, conv_state, ssm_state) -> Tuple[torch.Tensor, torch.Tensor]: + # Remove the chunked prefill request from the request / token counts so + # the subsequent prefill computation ignores the chunked prefill request. + active_token_count -= chunked_prefill_request_token_count + active_request_count -= 1 + else: + y_prefill_chunked = None + + # Process non-chunked prefill requests if there are any. + if (remaining_prefill_tokens := active_token_count - first_prefill_token_idx) > 0: + zxBCdt_prefill = zxBCdt[first_prefill_token_idx:active_token_count] + cu_seqlens_prefill = F.pad( + cu_seqlens[first_prefill_token_idx + 1 : active_request_count + 1] + - first_prefill_token_idx, + (1, 0), + ) + seq_idx_prefill = ( + seq_idx[:, first_prefill_token_idx:active_token_count] - first_prefill_token_idx + ) + batch_indices_prefill = batch_indices[first_prefill_token_idx:active_request_count] + + y_prefill = self.ssm_prefill( + zxBCdt_prefill, + conv_state=conv_state, + ssm_state=ssm_state, + seq_idx=seq_idx_prefill, + cu_seqlens=cu_seqlens_prefill, + return_varlen_states=return_varlen_states, + batch_indices=batch_indices_prefill, + ) + else: + y_prefill = None + + # Assemble the final output by concatenating the decode output, + # non-chunked prefill output, and chunked prefill output together. + y_prefill = maybe_cat(y_prefill, y_prefill_chunked, required=True) + y = maybe_cat(y_decode, y_prefill, required=True) + + # Add padding tokens back if necessary. Note that we use the context active token count + # in case we modified the local count for chunked prefill above. + if (num_padding_tokens := padded_active_token_count - context.active_token_count) > 0: + y = torch.cat((y, y.new_zeros(num_padding_tokens, *y.shape[1:])), dim=0) + + # The output projection will perform the sequence parallel reduce-scatter if necessary. + out, out_bias = self.out_proj(y) + + return out, out_bias + + def decode( + self, hidden_states, conv_state, ssm_state, batch_indices: Optional[torch.Tensor] = None + ) -> Tuple[torch.Tensor, torch.Tensor]: """Performs inference step for decoding.""" # assert self.ngroups_local_tp == 1, "Only support ngroups=1 for inference for now" - dtype = hidden_states.dtype - assert hidden_states.shape[0] == 1, "Only support decoding with 1 token at a time for now" + is_dynamic_batching = batch_indices is not None - # b d_model --> b p(2d) + if not is_dynamic_batching: + assert ( + hidden_states.shape[0] == 1 + ), "Only support decoding with 1 token at a time for now" + + # (1, b, d_model) -> (1, b, proj_dim) zxBCdt, _ = self.in_proj(hidden_states) + # Make batch size leading dimension since that is 1 + if is_dynamic_batching: + zxBCdt = zxBCdt.transpose(0, 1) + assert self.cp.cp_size == 1, "Context parallel not supported for Mamba inferenece decode" - y = self.ssm_decode(zxBCdt, conv_state=conv_state, ssm_state=ssm_state) + y = self.ssm_decode( + zxBCdt, conv_state=conv_state, ssm_state=ssm_state, batch_indices=batch_indices + ) + + # Restore sequence length as first dimension + if is_dynamic_batching: + y = y.transpose(0, 1) - # l b pd --> l b d + # y has shape (1, b, d_inner), which is what out_proj expects out, out_bias = self.out_proj(y) + return out, out_bias def ssm_training(self, zxBCdt: torch.Tensor) -> torch.Tensor: @@ -497,8 +626,34 @@ def ssm_prefill( zxBCdt: torch.Tensor, conv_state: Optional[torch.Tensor], ssm_state: Optional[torch.Tensor], + seq_idx: Optional[torch.Tensor] = None, + cu_seqlens: Optional[torch.Tensor] = None, + return_varlen_states: bool = False, + batch_indices: Optional[torch.Tensor] = None, + is_chunked_prefill: bool = False, ) -> torch.Tensor: - """Performs SSM computation for inference prefill step.""" + """ + Performs SSM computation for inference prefill step. + + Args: + zxBCdt: The input tensor of shape (l, b, d), which is a concatenation of + z, x, B, C, and dt projections. + conv_state: The convolution state tensor for inference. + ssm_state: The selective scan state tensor for inference. + seq_idx: A map from token index to request index for variable-length sequences. + cu_seqlens: Cumulative sequence lengths for variable-length sequences. + return_varlen_states: Whether to return variable-length states from the SSM kernel. + batch_indices: A map from batch id to position in the Mamba state tensors for + dynamic inference. + is_chunked_prefill: Whether the request is a chunked prefill request. + + Returns: + The output tensor of shape (l, b, d). + """ + is_dynamic_batching = seq_idx is not None + assert not ( + is_dynamic_batching and is_chunked_prefill + ), "Cannot use chunked prefill with dynamic batching" # transpose: l b pd --> b l pd zxBCdt = rearrange(zxBCdt, "l b d -> b l d").contiguous() @@ -516,29 +671,53 @@ def ssm_prefill( dim=-1, ) - # transpose: b l pd --> b pd l - xBC = rearrange(xBC, "b l d -> b d l").contiguous() - # Compute short convolution - if conv_state is not None: - # If we just take x[:, :, -self.d_conv :], it will error if seqlen < self.d_conv - # Instead F.pad will pad with zeros if seqlen < self.d_conv, and truncate otherwise. - conv_state.copy_(F.pad(xBC, (self.d_conv - xBC.shape[-1], 0))) # Update state (B D W) + if conv_state is not None and is_dynamic_batching: + # xBC should have shape (b l d) for causal_conv1d_varlen_states + assert batch_indices is not None + conv_state[batch_indices] = causal_conv1d_varlen_states( + xBC.squeeze(0), cu_seqlens, state_len=conv_state.shape[-1] + ) + + # Maintain channels-last memory layout to use seq_idx for causal_conv1d_fn + # See https://github.com/Dao-AILab/causal-conv1d/blob/69e6dadc28b169a4c49cb86b586f64ee90242c70/csrc/causal_conv1d.cpp#L174 # pylint: disable=line-too-long + xBC = xBC.transpose(1, 2) + elif is_chunked_prefill: + # Maintain channels-last memory layout to use initial_states for causal_conv1d_fn + # See https://github.com/Dao-AILab/causal-conv1d/blob/69e6dadc28b169a4c49cb86b586f64ee90242c70/csrc/causal_conv1d.cpp#L200 # pylint: disable=line-too-long + xBC = xBC.transpose(1, 2) + else: + # transpose: b l pd --> b pd l + xBC = rearrange(xBC, "b l d -> b d l").contiguous() + if conv_state is not None: + # If we just take x[:, :, -self.d_conv :], it will error if seqlen < self.d_conv + # Instead F.pad will pad with zeros if seqlen < self.d_conv, and truncate otherwise. + conv_state.copy_( + F.pad(xBC, (self.d_conv - xBC.shape[-1], 0)) + ) # Update state (B D W) seqlen = xBC.size(2) if causal_conv1d_fn is None: xBC = self.act(self.cp.conv1d(xBC)[..., :seqlen]) else: assert self.activation in ["silu", "swish"] + if is_chunked_prefill: + initial_conv_state = ( + conv_state[:, :, 1:].permute(0, 2, 1).contiguous().transpose(1, 2) + ) + else: + initial_conv_state = None xBC = causal_conv1d_fn( x=xBC, weight=rearrange(self.cp.get_conv1d_weight(), "d 1 w -> d w"), bias=self.cp.get_conv1d_bias(), activation=self.activation, + seq_idx=seq_idx, + initial_states=initial_conv_state, ) # transpose b pd l --> b l pd - xBC = rearrange(xBC, "b d l -> b l d").contiguous() + xBC = rearrange(xBC, "b d l -> b l d").contiguous() x, B, C = torch.split( xBC, @@ -565,6 +744,14 @@ def ssm_prefill( self.cp.cp_size == 1 or self.rmsnorm ), "Context parallel not supported for use_mem_eff_path==False and rmsnorm==False" + if is_chunked_prefill: + initial_ssm_state = ssm_state + else: + initial_ssm_state = None + + # Note that both `seq_idx` and `cu_seqlens` must be passed in + # for variable length generation. + # See https://github.com/state-spaces/mamba/blob/e0761ece1db07e0949dd88b4f4cd440420a19fd9/tests/test_generation.py#L97 # pylint: disable=line-too-long y = mamba_chunk_scan_combined( x, dt, @@ -581,11 +768,25 @@ def ssm_prefill( dt_bias=self.cp.get_dt_bias().float(), dt_softplus=True, return_final_states=ssm_state is not None, + seq_idx=seq_idx, + cu_seqlens=cu_seqlens, + return_varlen_states=return_varlen_states, + initial_states=initial_ssm_state, ) if ssm_state is not None: - y, last_state = y - ssm_state.copy_(last_state) + if return_varlen_states: + assert batch_indices is not None + + y, _, varlen_states = y + + # This has to be varlen_states, NOT last_state + # See reference implementation: + # https://github.com/state-spaces/mamba/blob/e0761ece1db07e0949dd88b4f4cd440420a19fd9/mamba_ssm/modules/mamba2.py#L267 # pylint: disable=line-too-long + ssm_state[batch_indices] = varlen_states + else: + y, last_state = y + ssm_state.copy_(last_state) y = rearrange(y, "b l h p -> l b (h p)").contiguous() y = self.cp.post_conv_ssm(y) @@ -598,14 +799,31 @@ def ssm_prefill( return y def ssm_decode( - self, zxBCdt: torch.Tensor, conv_state: torch.Tensor, ssm_state: torch.Tensor + self, + zxBCdt: torch.Tensor, + conv_state: torch.Tensor, + ssm_state: torch.Tensor, + batch_indices: Optional[torch.Tensor] = None, ) -> torch.Tensor: - """Performs SSM computation for inference decode step.""" - + """ + Performs SSM computation for inference decode step. + + Args: + zxBCdt: The input tensor of shape (l, b, d), which is a concatenation of + z, x, B, C, and dt projections. For decoding, l must be 1. + conv_state: The convolution state tensor for inference. + ssm_state: The selective scan state tensor for inference. + batch_indices: A map from batch id to position in the Mamba state tensors for + dynamic inference. + + Returns: + The output tensor of shape (l, b, d). + """ + seq_len, batch_size, _ = zxBCdt.shape dtype = zxBCdt.dtype - assert zxBCdt.shape[0] == 1, "Only support decoding with 1 token at a time for now" + assert seq_len == 1, "Only support decoding with 1 token at a time for now" - # l b d --> b d + # Remove sequence dimension zxBCdt = zxBCdt.squeeze(0) z, xBC, dt = torch.split( @@ -627,7 +845,7 @@ def ssm_decode( ) # (B D) if self.conv1d.bias is not None: xBC = xBC + self.conv1d.bias - xBC = self.act(xBC).to(dtype=dtype) + xBC = self.act(xBC).to(dtype=xBC.dtype) else: xBC = causal_conv1d_update( xBC, @@ -635,6 +853,7 @@ def ssm_decode( rearrange(self.conv1d.weight, "d 1 w -> d w"), self.conv1d.bias, self.activation, + conv_state_indices=batch_indices, ) x, B, C = torch.split( @@ -715,35 +934,61 @@ def ssm_decode( z=z if not self.rmsnorm else None, dt_bias=dt_bias, dt_softplus=True, + state_batch_indices=batch_indices, ) y = rearrange(y, "b h p -> b (h p)") if self.rmsnorm: y = self.norm(y, z) - # b (h p) -> l b (h p) + # Restore sequence dimension return y.unsqueeze(0) - def allocate_inference_cache(self, batch_size, max_seqlen, dtype=None): - """ - allocate inference cache + def _get_varlen_generation_state( + self, inference_context: Optional[BaseInferenceContext] = None + ) -> Tuple[torch.Tensor, torch.Tensor, bool]: + """Constructs the variable length generation state for non-decode dynamic inference. + + The returned state includes the following: + `seq_idx` (Tensor): A map from token idx to request idx. + `cu_seqlens` (Tensor): The cumulative sequence lengths. + `return_varlen_states` (bool): Whether to return a varlen states tensor for + `mamba_chunk_scan_combined`. + + Returns empty state for training, static inference, or decode-only dynamic inference. + + Args: + inference_context (InferenceContext): The inference context. + + Returns: + A tuple of (`seq_idx`, `cu_seqlens`, `return_varlen_states`) """ - device = self.out_proj.weight.device - conv_dtype = self.conv1d.weight.dtype if dtype is None else dtype - conv_state = torch.zeros( - batch_size, self.conv1d.weight.shape[0], self.d_conv, device=device, dtype=conv_dtype - ) - ssm_dtype = self.in_proj.weight.dtype if dtype is None else dtype - # ssm_dtype = torch.float32 - ssm_state = torch.zeros( - batch_size, - self.nheads_local_tp, - self.headdim, - self.d_state, - device=device, - dtype=ssm_dtype, + + if ( + inference_context is None + or not inference_context.is_dynamic_batching() + or inference_context.is_decode_only() + ): + return None, None, False + + active_token_count = inference_context.active_token_count + seq_idx = ( + inference_context.token_to_request_idx[:active_token_count] + .clone() + .to(torch.int32) + .unsqueeze(0) ) - return conv_state, ssm_state + + # Get the list of cumulative sequence lengths for active requests. + cu_seqlens, _ = inference_context.cu_query_lengths() + + return seq_idx, cu_seqlens, True + + def mamba_state_shapes_per_request(self) -> Tuple[Tuple[int], Tuple[int]]: + """Returns the Mamba conv and ssm states shapes per request.""" + conv_states_shape = (self.conv1d.weight.shape[0], self.d_conv) + ssm_states_shape = (self.nheads_local_tp, self.headdim, self.d_state) + return (conv_states_shape, ssm_states_shape) def _get_states_from_cache(self, inference_context, batch_size, *, inference_params=None): """Initializes or retrieves the SSM state tensors from the cache. @@ -756,23 +1001,23 @@ def _get_states_from_cache(self, inference_context, batch_size, *, inference_par inference_context = deprecate_inference_params(inference_context, inference_params) assert inference_context is not None + assert inference_context.is_static_batching() assert self.layer_number is not None + if ( self.layer_number not in inference_context.key_value_memory_dict or batch_size != self.cached_batch_size ): + conv_state_shape, ssm_state_shape = self.mamba_state_shapes_per_request() conv_state = torch.zeros( batch_size, - self.conv1d.weight.shape[0], - self.d_conv, + *conv_state_shape, device=self.conv1d.weight.device, dtype=self.conv1d.weight.dtype, ) ssm_state = torch.zeros( batch_size, - self.nheads_local_tp, - self.headdim, - self.d_state, + *ssm_state_shape, device=self.in_proj.weight.device, dtype=self.in_proj.weight.dtype, ) @@ -780,7 +1025,6 @@ def _get_states_from_cache(self, inference_context, batch_size, *, inference_par self.cached_batch_size = batch_size else: conv_state, ssm_state = inference_context.key_value_memory_dict[self.layer_number] - # TODO: Remove reference to `inference_context.sequence_len_offset` for dynamic batching if inference_context.sequence_len_offset == 0: conv_state.zero_() ssm_state.zero_() diff --git a/megatron/core/utils.py b/megatron/core/utils.py index abfaf7f6320..93b2e593d84 100644 --- a/megatron/core/utils.py +++ b/megatron/core/utils.py @@ -65,6 +65,8 @@ _torch_version = PkgVersion("0.0.0") if HAVE_PACKAGING else "0.0.0" _te_version = None _fa_version = None +_mamba_ssm_version = None +_causal_conv1d_version = None @contextmanager @@ -388,6 +390,79 @@ def is_fa_min_version(version, check_equality=True): return get_fa_version() > PkgVersion(version) +def get_mamba_version(): + """Get mamba version from __version__; if not available use pip's. Use caching.""" + if not HAVE_PACKAGING: + raise ImportError( + "packaging is not installed. Please install it with `pip install packaging`." + ) + + def get_mamba_version_str(): + import mamba_ssm + + if hasattr(mamba_ssm, "__version__"): + return str(mamba_ssm.__version__) + else: + return version("mamba_ssm") + + global _mamba_ssm_version + if _mamba_ssm_version is None: + _mamba_ssm_version = PkgVersion(get_mamba_version_str()) + return _mamba_ssm_version + + +def is_mamba_min_version(version, check_equality=True): + """Check if minimum version of `mamba_ssm` is installed.""" + if not HAVE_PACKAGING: + raise ImportError( + "packaging is not installed. Please install it with `pip install packaging`." + ) + if check_equality: + return get_mamba_version() >= PkgVersion(version) + return get_mamba_version() > PkgVersion(version) + + +def get_causal_conv1d_version(): + """Get causal_conv1d version from __version__; if not available use pip's. Use caching.""" + if not HAVE_PACKAGING: + raise ImportError( + "packaging is not installed. Please install it with `pip install packaging`." + ) + + def get_causal_conv1d_version_str(): + import causal_conv1d + + if hasattr(causal_conv1d, "__version__"): + return str(causal_conv1d.__version__) + else: + return version("causal_conv1d") + + global _causal_conv1d_version + if _causal_conv1d_version is None: + _causal_conv1d_version = PkgVersion(get_causal_conv1d_version_str()) + return _causal_conv1d_version + + +def is_causal_conv1d_min_version(version, check_equality=True): + """Check if minimum version of `causal_conv1d` is installed.""" + if not HAVE_PACKAGING: + raise ImportError( + "packaging is not installed. Please install it with `pip install packaging`." + ) + if check_equality: + return get_causal_conv1d_version() >= PkgVersion(version) + return get_causal_conv1d_version() > PkgVersion(version) + + +def check_mamba_sequence_packing_support() -> Tuple[bool, Optional[str]]: + """Checks whether `causal_conv1d` and `mamba_ssm` support sequence packing.""" + if not is_causal_conv1d_min_version("1.5.3.post1"): + return False, "causal_conv1d >= 1.5.3.post1 is required" + elif not is_mamba_min_version("2.2.6.post3"): + return False, "mamba_ssm >= 2.2.6.post3 is required" + return True, None + + def ensure_divisibility(numerator, denominator): """Ensure that numerator is divisible by the denominator.""" assert numerator % denominator == 0, "{} is not divisible by {}".format(numerator, denominator) @@ -2001,6 +2076,16 @@ def unwrap_model(model, module_instances=None): return unwrapped_model +def maybe_cat(a, b, dim=0, *, required=False): + """Concatenates `a` and `b` along `dim` if `a` and `b` exist.""" + xs = [t for t in (a, b) if t is not None] + if not xs: + if required: + raise ValueError("both tensors are None") + return None + return xs[0] if len(xs) == 1 else torch.cat(xs, dim=dim) + + def get_asyncio_loop(loop: asyncio.AbstractEventLoop | None = None) -> asyncio.AbstractEventLoop: """Creates an asyncio loop if necessary and then returns the current asyncio loop.""" if loop is None: diff --git a/megatron/training/tokenizer/sft_tokenizer.py b/megatron/training/tokenizer/sft_tokenizer.py index 4a941fc180b..f525352e892 100644 --- a/megatron/training/tokenizer/sft_tokenizer.py +++ b/megatron/training/tokenizer/sft_tokenizer.py @@ -170,6 +170,11 @@ def pad(self): """Pad token ID.""" return self._prompt_config.pad_token_id + @property + def bos(self): + """Beginning of sequence token ID.""" + return self._tokenizer.bos_token_id + @property def eod(self): """End of sentence token ID.""" diff --git a/tests/unit_tests/inference/contexts/test_dynamic_context.py b/tests/unit_tests/inference/contexts/test_dynamic_context.py index 1cd9d66ece1..0674cdfcabd 100644 --- a/tests/unit_tests/inference/contexts/test_dynamic_context.py +++ b/tests/unit_tests/inference/contexts/test_dynamic_context.py @@ -12,6 +12,7 @@ ) from megatron.core.inference.inference_request import DynamicInferenceRequest from megatron.core.inference.sampling_params import SamplingParams +from megatron.core.ssm.mamba_hybrid_layer_allocation import Symbols from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from tests.unit_tests.test_utilities import Utils @@ -42,12 +43,19 @@ def _get_dynamic_context( max_sequence_length, buffer_size_gb, block_size_tokens, - buffer_guarenteed_fraction, + buffer_guaranteed_fraction, buffer_overflow_factor, max_requests_override, max_tokens_override, + is_hybrid_model=False, + layer_type_list=None, + rounder=64, ): - set_rounder(64) + set_rounder(rounder) + + if is_hybrid_model and layer_type_list is None: + layer_type_list = [Symbols.MAMBA, Symbols.MLP, Symbols.ATTENTION, Symbols.MLP] + dynamic_context = DynamicInferenceContext( params_dtype=params_dtype, num_layers=num_layers, @@ -55,23 +63,27 @@ def _get_dynamic_context( num_attention_heads=num_attention_heads, max_sequence_length=max_sequence_length, num_cuda_graphs=None, + use_cuda_graphs_for_non_decode_steps=not is_hybrid_model, buffer_size_gb=buffer_size_gb, - buffer_guaranteed_fraction=buffer_guarenteed_fraction, + buffer_guaranteed_fraction=buffer_guaranteed_fraction, block_size_tokens=block_size_tokens, buffer_overflow_factor=buffer_overflow_factor, max_requests_override=max_requests_override, max_tokens_override=max_tokens_override, + layer_type_list=layer_type_list, + mamba_conv_states_shape=(544, 4), + mamba_ssm_states_shape=(8, 64, 16), use_flashinfer_fused_rope=None, # default to using flash-infer if available # this is for compatibility with the LTS environment ) return dynamic_context def teardown_method(self, method): - set_rounder(64) Utils.destroy_model_parallel() @pytest.mark.internal - def test_initialize_dynamic_context(self): + @pytest.mark.parametrize("is_hybrid_model", [False, True]) + def test_initialize_dynamic_context(self, is_hybrid_model: bool): self._setup_model_parallel_group(1, 1) dynamic_context = self._get_dynamic_context( @@ -81,18 +93,30 @@ def test_initialize_dynamic_context(self): num_attention_heads=2, max_sequence_length=512, buffer_size_gb=0.03, - buffer_guarenteed_fraction=0.1, + buffer_guaranteed_fraction=0.1, block_size_tokens=128, max_requests_override=None, max_tokens_override=None, buffer_overflow_factor=None, + is_hybrid_model=is_hybrid_model, ) - assert dynamic_context.gtd_block_count == 48 - assert dynamic_context.gtd_request_count == 12 - assert dynamic_context.block_allocator.block_count_total == 491 - assert dynamic_context.max_requests == 128 - assert dynamic_context.max_tokens == 62848 + if not is_hybrid_model: + assert dynamic_context.gtd_block_count == 48 + assert dynamic_context.gtd_request_count == 12 + assert dynamic_context.block_allocator.block_count_total == 491 + assert dynamic_context.max_requests == 128 + assert dynamic_context.max_tokens == 62848 + assert dynamic_context.num_mamba_layers == 0 + assert dynamic_context.mamba_metadata is None + else: + assert dynamic_context.gtd_block_count == 112 + assert dynamic_context.gtd_request_count == 28 + assert dynamic_context.block_allocator.block_count_total == 1156 + assert dynamic_context.max_requests == 320 + assert dynamic_context.max_tokens == 154176 + assert dynamic_context.num_mamba_layers == 1 + assert dynamic_context.mamba_metadata is not None # Check initializations to -1 assert torch.all(dynamic_context.request_ids == -1) @@ -100,32 +124,38 @@ def test_initialize_dynamic_context(self): @pytest.mark.internal def test_is_static_batching(self): self._setup_model_parallel_group(1, 1) - dynamic_context = DynamicInferenceContext( + dynamic_context = self._get_dynamic_context( params_dtype=torch.float32, num_layers=2, kv_channels=64, num_attention_heads=8, max_sequence_length=512, - num_cuda_graphs=None, buffer_size_gb=1.0, buffer_guaranteed_fraction=0.1, block_size_tokens=128, + max_requests_override=None, + max_tokens_override=None, + buffer_overflow_factor=None, ) assert not dynamic_context.is_static_batching() @pytest.mark.internal - def test_is_memory_available(self): + @pytest.mark.parametrize("is_hybrid_model", [False, True]) + def test_is_memory_available(self, is_hybrid_model): self._setup_model_parallel_group(1, 1) - dynamic_context = DynamicInferenceContext( + dynamic_context = self._get_dynamic_context( params_dtype=torch.float32, num_layers=2, kv_channels=64, num_attention_heads=8, max_sequence_length=512, - num_cuda_graphs=None, buffer_size_gb=1.0, buffer_guaranteed_fraction=0.1, block_size_tokens=128, + max_requests_override=None, + max_tokens_override=None, + buffer_overflow_factor=None, + is_hybrid_model=is_hybrid_model, ) dynamic_context.block_allocator.block_count_avail = 10 assert dynamic_context.block_allocator.is_memory_available(10) @@ -141,19 +171,24 @@ def test_is_memory_available(self): assert not dynamic_context.block_allocator.is_memory_available(6, safe=True) @pytest.mark.internal - def test_request_overflow(self): + @pytest.mark.parametrize("is_hybrid_model", [False, True]) + def test_request_overflow(self, is_hybrid_model: bool): self._setup_model_parallel_group(1, 1) - set_rounder(1) - dynamic_context = DynamicInferenceContext( + + dynamic_context = self._get_dynamic_context( params_dtype=torch.float32, num_layers=2, kv_channels=64, num_attention_heads=8, max_sequence_length=128, - num_cuda_graphs=None, buffer_size_gb=0.01, buffer_guaranteed_fraction=0.1, block_size_tokens=32, + max_requests_override=None, + max_tokens_override=None, + buffer_overflow_factor=None, + rounder=1, + is_hybrid_model=is_hybrid_model, ) with pytest.raises(RequestOverflowError): for i in range(dynamic_context.max_requests + 1): @@ -168,22 +203,24 @@ def test_request_overflow(self): ) # Adding more than allowed requests @pytest.mark.internal - def test_token_overflow_error(self): + @pytest.mark.parametrize("is_hybrid_model", [False, True]) + def test_token_overflow_error(self, is_hybrid_model: bool): self._setup_model_parallel_group(1, 1) - set_rounder(1) - dynamic_context = DynamicInferenceContext( + + dynamic_context = self._get_dynamic_context( params_dtype=torch.float32, num_layers=2, kv_channels=64, num_attention_heads=8, max_sequence_length=512, - num_cuda_graphs=None, buffer_size_gb=0.1, buffer_guaranteed_fraction=0.1, block_size_tokens=128, buffer_overflow_factor=1.0, max_requests_override=2, max_tokens_override=20, # Setting a very low token limit + rounder=1, + is_hybrid_model=is_hybrid_model, ) with pytest.raises(TokenOverflowError): @@ -198,18 +235,23 @@ def test_token_overflow_error(self): ) # Exceeding max token count @pytest.mark.internal - def test_reset(self): + @pytest.mark.parametrize("is_hybrid_model", [False, True]) + def test_reset(self, is_hybrid_model: bool): self._setup_model_parallel_group(1, 1) - dynamic_context = DynamicInferenceContext( + + dynamic_context = self._get_dynamic_context( params_dtype=torch.float32, num_layers=2, kv_channels=64, num_attention_heads=8, max_sequence_length=128, - num_cuda_graphs=None, buffer_size_gb=1.0, buffer_guaranteed_fraction=0.1, block_size_tokens=128, + max_requests_override=None, + max_tokens_override=None, + buffer_overflow_factor=None, + is_hybrid_model=is_hybrid_model, ) # Initialize all variables @@ -234,6 +276,9 @@ def test_reset(self): dynamic_context.block_allocator.block_count_avail = 5 dynamic_context.memory_buffer.fill_(1) dynamic_context.request_to_kv_block_ids.fill_(1) + if is_hybrid_model: + dynamic_context.mamba_conv_states.fill_(1) + dynamic_context.mamba_ssm_states.fill_(1) # Call reset dynamic_context.reset() @@ -262,9 +307,14 @@ def test_reset(self): == dynamic_context.block_allocator.block_count_total - 1 ) assert torch.all(dynamic_context.request_to_kv_block_ids == -1) + if is_hybrid_model: + assert torch.all(dynamic_context.mamba_metadata.request_to_mamba_state_idx == -1) + assert torch.all(dynamic_context.mamba_conv_states == 0) + assert torch.all(dynamic_context.mamba_ssm_states == 0) @pytest.mark.internal - def test_allocate_and_release_memory_blocks(self): + @pytest.mark.parametrize("is_hybrid_model", [False, True]) + def test_allocate_and_release_memory_blocks(self, is_hybrid_model): self._setup_model_parallel_group(1, 1) dynamic_context = self._get_dynamic_context( params_dtype=torch.float32, @@ -273,23 +323,38 @@ def test_allocate_and_release_memory_blocks(self): num_attention_heads=2, max_sequence_length=512, buffer_size_gb=0.03, - buffer_guarenteed_fraction=0.1, + buffer_guaranteed_fraction=0.1, block_size_tokens=128, max_requests_override=None, max_tokens_override=None, buffer_overflow_factor=None, + is_hybrid_model=is_hybrid_model, ) - assert dynamic_context.block_allocator.allocate_memory_blocks( - 4 - ).cpu().detach().numpy().tolist() == [486, 487, 488, 489] - assert dynamic_context.block_allocator.block_count_avail == 486 + if is_hybrid_model: + expected_memory_blocks = [1151, 1152, 1153, 1154] + else: + expected_memory_blocks = [486, 487, 488, 489] + expected_block_count_avail = expected_memory_blocks[0] + + assert ( + dynamic_context.block_allocator.allocate_memory_blocks(4) + .cpu() + .detach() + .numpy() + .tolist() + == expected_memory_blocks + ) + assert dynamic_context.block_allocator.block_count_avail == expected_block_count_avail dynamic_context.block_allocator.release_memory_blocks( - torch.tensor([488, 489], device='cuda') + torch.tensor(expected_memory_blocks[-2:], device='cuda') ) - assert dynamic_context.block_allocator.block_count_avail == 488 - assert dynamic_context.block_allocator.allocate_memory_blocks(1).item() == 489 - assert dynamic_context.block_allocator.block_count_avail == 487 + assert dynamic_context.block_allocator.block_count_avail == expected_block_count_avail + 2 + assert ( + dynamic_context.block_allocator.allocate_memory_blocks(1).item() + == expected_memory_blocks[-1] + ) + assert dynamic_context.block_allocator.block_count_avail == expected_block_count_avail + 1 # Should return None since we allocate more blocks than what we have. assert ( dynamic_context.block_allocator.allocate_memory_blocks( @@ -299,8 +364,10 @@ def test_allocate_and_release_memory_blocks(self): ) @pytest.mark.internal - def test_add_request(self): + @pytest.mark.parametrize("is_hybrid_model", [False, True]) + def test_add_request(self, is_hybrid_model: bool): self._setup_model_parallel_group(1, 1) + dynamic_context = self._get_dynamic_context( params_dtype=torch.float32, num_layers=4, @@ -308,11 +375,12 @@ def test_add_request(self): num_attention_heads=2, max_sequence_length=512, buffer_size_gb=0.03, - buffer_guarenteed_fraction=0.1, + buffer_guaranteed_fraction=0.1, block_size_tokens=128, max_requests_override=None, max_tokens_override=None, buffer_overflow_factor=None, + is_hybrid_model=is_hybrid_model, ) assert dynamic_context.block_size_tokens == 128 context_length = 144 @@ -331,14 +399,10 @@ def test_add_request(self): assert torch.all(dynamic_context.request_ids[1:] == -1) assert dynamic_context.request_query_lengths[0] == context_length assert dynamic_context.request_kv_length_offsets[0] == 0 - assert dynamic_context.request_to_kv_block_ids[0].cpu().detach().numpy().tolist() == [ - 488, - 489, - -1, - -1, - ] assert dynamic_context.request_kv_block_counts[0] == 2 - assert dynamic_context.request_last_kv_block_id[0] == 489 + assert dynamic_context.request_last_kv_block_id[0].item() == ( + 1154 if is_hybrid_model else 489 + ) assert dynamic_context.request_last_kv_block_offset[0].item() == 15 assert torch.all( dynamic_context.token_to_pos_ids[0:context_length] @@ -352,17 +416,22 @@ def test_add_request(self): dynamic_context.token_to_position_in_request[0:context_length] == torch.arange(0, context_length, dtype=torch.long, device='cuda') ) + + # Verify token_to_block_idx and token_to_local_position_within_kv_block based on assigned blocks + first_block_id = dynamic_context.request_to_kv_block_ids[0, 0] + second_block_id = dynamic_context.request_to_kv_block_ids[0, 1] + assert torch.all( dynamic_context.token_to_block_idx[0:context_length][ 0 : dynamic_context.block_size_tokens ] - == 488 + == first_block_id ) assert torch.all( dynamic_context.token_to_block_idx[0:context_length][ dynamic_context.block_size_tokens : context_length ] - == 489 + == second_block_id ) assert torch.all( dynamic_context.token_to_local_position_within_kv_block[0:context_length] @@ -371,8 +440,10 @@ def test_add_request(self): ) @pytest.mark.internal - def test_update_request(self): + @pytest.mark.parametrize("is_hybrid_model", [False, True]) + def test_update_request(self, is_hybrid_model: bool): self._setup_model_parallel_group(1, 1) + dynamic_context = self._get_dynamic_context( params_dtype=torch.float32, num_layers=4, @@ -380,11 +451,12 @@ def test_update_request(self): num_attention_heads=2, max_sequence_length=512, buffer_size_gb=0.03, - buffer_guarenteed_fraction=0.1, + buffer_guaranteed_fraction=0.1, block_size_tokens=128, max_requests_override=None, max_tokens_override=None, buffer_overflow_factor=None, + is_hybrid_model=is_hybrid_model, ) # This case should just reset and return since all requests are finished @@ -394,10 +466,19 @@ def test_update_request(self): dynamic_context.request_kv_block_counts[0:3] = 1 new_block_ids = dynamic_context.block_allocator.allocate_memory_blocks(3, safe=True) dynamic_context.request_to_kv_block_ids[0:3, 0] = new_block_ids + + if is_hybrid_model: + # Also initialize Mamba states for the dummy requests + dynamic_context.mamba_conv_states[:, 0:3, :, :].fill_(1.0) + dynamic_context.mamba_ssm_states[:, 0:3, :, :, :].fill_(1.0) + dynamic_context.update_requests( active_requests_mask=active_requests_mask, new_tokens=torch.tensor([0, 1, 2]) ) assert dynamic_context.total_request_count == 0 + if is_hybrid_model: + assert torch.all(dynamic_context.mamba_conv_states == 0) + assert torch.all(dynamic_context.mamba_ssm_states == 0) # This case would cover all cases # 1. Already there will be 2 paused requests @@ -406,9 +487,9 @@ def test_update_request(self): # 4. Some of these requests will be resumed. # Setup is as follows : # Request ids 0, 1 are paused - # Request ids 2 , 4, 9 are active requests + # Request ids 2, 4, 9 are active requests # Request ids 3 7 8 have completed - # Request ids 5 and 6 will require on more block later on coz they finished their current block + # Request ids 5 and 6 will require on more block later on because they finished their current block dynamic_context = self._get_dynamic_context( params_dtype=torch.float32, @@ -417,11 +498,12 @@ def test_update_request(self): num_attention_heads=2, max_sequence_length=512, buffer_size_gb=0.03, - buffer_guarenteed_fraction=0.1, + buffer_guaranteed_fraction=0.1, block_size_tokens=128, max_requests_override=None, max_tokens_override=None, buffer_overflow_factor=None, + is_hybrid_model=is_hybrid_model, ) active_requests_mask = torch.Tensor([1, 0, 1, 1, 1, 0, 0, 1]).cuda().int() @@ -472,6 +554,14 @@ def test_update_request(self): dynamic_context.request_last_kv_block_offset[0:2] = dynamic_context.block_size_tokens - 1 dynamic_context.request_last_kv_block_offset[5:7] = dynamic_context.block_size_tokens - 1 + if is_hybrid_model: + # Dummy fill for states to be non-zero before update + for i in range(total_request_count): + dynamic_context.mamba_metadata.request_to_mamba_state_idx[i] = i + dynamic_context.mamba_metadata.mamba_state_free_slot_count -= total_request_count + dynamic_context.mamba_conv_states[:, 0:total_request_count, :, :] = 1.0 + dynamic_context.mamba_ssm_states[:, 0:total_request_count, :, :, :] = 1.0 + dynamic_context.update_requests( active_requests_mask=active_requests_mask, new_tokens=next_tokens ) @@ -522,28 +612,49 @@ def test_update_request(self): # The first 4 requests will require an extra block. # Since 3 requests have finished, the last 3 rows should be all -1. - assert torch.all( - dynamic_context.request_to_kv_block_ids[0:10].cpu() - == torch.tensor( - [ - [479, 482, -1, -1], - [480, 479, -1, -1], - [484, 486, -1, -1], - [485, 487, -1, -1], - [483, -1, -1, -1], - [481, -1, -1, -1], - [488, -1, -1, -1], - [-1, -1, -1, -1], - [-1, -1, -1, -1], - [-1, -1, -1, -1], - ] + if is_hybrid_model: + assert torch.all( + dynamic_context.request_to_kv_block_ids[0:10].cpu() + == torch.tensor( + [ + [1144, 1147, -1, -1], + [1145, 1144, -1, -1], + [1149, 1151, -1, -1], + [1150, 1152, -1, -1], + [1148, -1, -1, -1], + [1146, -1, -1, -1], + [1153, -1, -1, -1], + [-1, -1, -1, -1], + [-1, -1, -1, -1], + [-1, -1, -1, -1], + ] + ) + ) + else: + assert torch.all( + dynamic_context.request_to_kv_block_ids[0:10].cpu() + == torch.tensor( + [ + [479, 482, -1, -1], + [480, 479, -1, -1], + [484, 486, -1, -1], + [485, 487, -1, -1], + [483, -1, -1, -1], + [481, -1, -1, -1], + [488, -1, -1, -1], + [-1, -1, -1, -1], + [-1, -1, -1, -1], + [-1, -1, -1, -1], + ] + ) ) - ) @pytest.mark.internal - def test_release_memory_blocks_for_finished_requests(self): + @pytest.mark.parametrize("is_hybrid_model", [False, True]) + def test_release_memory_blocks_for_finished_requests(self, is_hybrid_model): """Test that memory blocks are correctly released for finished requests.""" self._setup_model_parallel_group(1, 1) + dynamic_context = self._get_dynamic_context( params_dtype=torch.float32, num_layers=4, @@ -551,11 +662,12 @@ def test_release_memory_blocks_for_finished_requests(self): num_attention_heads=2, max_sequence_length=512, buffer_size_gb=0.03, - buffer_guarenteed_fraction=0.1, + buffer_guaranteed_fraction=0.1, block_size_tokens=128, max_requests_override=None, max_tokens_override=None, buffer_overflow_factor=None, + is_hybrid_model=is_hybrid_model, ) # Set up the initial state with 5 requests @@ -572,6 +684,13 @@ def test_release_memory_blocks_for_finished_requests(self): dynamic_context.request_to_kv_block_ids[i, 0] = initial_blocks[i] dynamic_context.request_query_lengths[i] = 1 dynamic_context.request_ids[i] = i + if is_hybrid_model: + dynamic_context.mamba_conv_states[:, i, :, :].fill_( + float(i + 1) + ) # Fill with distinct values + dynamic_context.mamba_ssm_states[:, i, :, :, :].fill_(float(i + 1)) + dynamic_context.mamba_metadata.request_to_mamba_state_idx[i] = i + dynamic_context.mamba_metadata.mamba_state_free_slot_count -= 1 # Create an active_requests_mask where requests 0, 2, and 4 are finished (0), # and requests 1 and 3 are still active (1) @@ -591,10 +710,26 @@ def test_release_memory_blocks_for_finished_requests(self): # Verify that 3 blocks were released by checking the available blocks assert dynamic_context.block_allocator.block_count_avail == initial_available_blocks + 3 + if is_hybrid_model: + # Request at position 3 now moves into finished request position 0 + # Request at position 1 remains active + mamba_idx = { + i: dynamic_context.mamba_metadata.request_to_mamba_state_idx[i] for i in range(5) + } + assert torch.all(dynamic_context.mamba_conv_states[:, mamba_idx[0], :, :] == 4.0) + assert torch.all(dynamic_context.mamba_ssm_states[:, mamba_idx[0], :, :, :] == 4.0) + assert torch.all(dynamic_context.mamba_conv_states[:, mamba_idx[1], :, :] == 2.0) + assert torch.all(dynamic_context.mamba_ssm_states[:, mamba_idx[1], :, :, :] == 2.0) + assert mamba_idx[2] == -1 + assert mamba_idx[3] == -1 + assert mamba_idx[4] == -1 + @pytest.mark.internal - def test_finished_requests_with_multiple_blocks(self): + @pytest.mark.parametrize("is_hybrid_model", [False, True]) + def test_finished_requests_with_multiple_blocks(self, is_hybrid_model): """Test that all memory blocks are correctly released for finished requests that use multiple blocks.""" self._setup_model_parallel_group(1, 1) + dynamic_context = self._get_dynamic_context( params_dtype=torch.float32, num_layers=4, @@ -602,11 +737,12 @@ def test_finished_requests_with_multiple_blocks(self): num_attention_heads=2, max_sequence_length=512, buffer_size_gb=0.03, - buffer_guarenteed_fraction=0.1, + buffer_guaranteed_fraction=0.1, block_size_tokens=128, max_requests_override=None, max_tokens_override=None, buffer_overflow_factor=None, + is_hybrid_model=is_hybrid_model, ) # Set up the initial state with 3 requests, where some use multiple blocks @@ -638,6 +774,9 @@ def test_finished_requests_with_multiple_blocks(self): for i in range(3): dynamic_context.request_query_lengths[i] = 1 dynamic_context.request_ids[i] = i + if is_hybrid_model: + dynamic_context.mamba_conv_states[:, i, :, :].fill_(float(i + 1)) + dynamic_context.mamba_ssm_states[:, i, :, :, :].fill_(float(i + 1)) # Create an active_requests_mask where all requests are finished active_requests_mask = torch.tensor([0, 0, 0], device=torch.cuda.current_device()) @@ -655,6 +794,92 @@ def test_finished_requests_with_multiple_blocks(self): # Verify that all 6 blocks were released by checking the available blocks assert dynamic_context.block_allocator.block_count_avail == initial_available_blocks + 6 + if is_hybrid_model: + # All mamba states should be zeroed out + assert torch.all(dynamic_context.mamba_conv_states == 0) + assert torch.all(dynamic_context.mamba_ssm_states == 0) + + @pytest.mark.internal + @pytest.mark.parametrize("is_hybrid_model", [False, True]) + def test_mamba_states_cache(self, is_hybrid_model: bool): + self._setup_model_parallel_group(1, 1) + + if not is_hybrid_model: + # If not hybrid, mamba_states_cache should fail + dynamic_context = self._get_dynamic_context( + params_dtype=torch.float32, + num_layers=4, + kv_channels=8, + num_attention_heads=2, + max_sequence_length=512, + buffer_size_gb=0.03, + buffer_guaranteed_fraction=0.1, + block_size_tokens=128, + max_requests_override=None, + max_tokens_override=None, + buffer_overflow_factor=None, + is_hybrid_model=False, + ) + with pytest.raises(AssertionError) as error: + conv_state, ssm_state = dynamic_context.mamba_states_cache(layer_number=1) + return + + dynamic_context = self._get_dynamic_context( + params_dtype=torch.float32, + num_layers=4, + kv_channels=8, + num_attention_heads=2, + max_sequence_length=512, + buffer_size_gb=0.03, + buffer_guaranteed_fraction=0.1, + block_size_tokens=128, + max_requests_override=None, + max_tokens_override=None, + buffer_overflow_factor=None, + is_hybrid_model=is_hybrid_model, + layer_type_list=[Symbols.MAMBA, Symbols.ATTENTION, Symbols.MAMBA, Symbols.ATTENTION], + ) + + # Add a request to populate states + context_length = 10 + dynamic_context.add_request( + DynamicInferenceRequest( + request_id=0, + prompt_tokens=torch.arange(0, context_length, dtype=torch.long, device='cuda'), + sampling_params=SamplingParams( + num_tokens_to_generate=dynamic_context.max_tokens - 10 + ), + ) + ) + dynamic_context.initialize_attention_state() + + # Manually set some dummy values in mamba_conv_states and mamba_ssm_states + # Mamba layers are at global indices 0 and 2 (mapped to local 0 and 1 via layer_map) + # `layer_map` will map global layer index to the corresponding Mamba/Attention index. + # For layer_type_list ["MAMBA", "ATTENTION", "MAMBA", "ATTENTION"], + # global layer 1 (index 0) is MAMBA -> local mamba layer 0 + # global layer 3 (index 2) is MAMBA -> local mamba layer 1 + + # Test for the first Mamba layer (global layer 1, local mamba layer 0) + global_layer_1_mamba_local_idx = 0 + dynamic_context.mamba_conv_states[global_layer_1_mamba_local_idx] = 10.0 + dynamic_context.mamba_ssm_states[global_layer_1_mamba_local_idx] = 20.0 + + # Test for the second Mamba layer (global layer 3, local mamba layer 1) + global_layer_3_mamba_local_idx = 1 + dynamic_context.mamba_conv_states[global_layer_3_mamba_local_idx] = 30.0 + dynamic_context.mamba_ssm_states[global_layer_3_mamba_local_idx] = 40.0 + + # Retrieve states using mamba_states_cache for global layer 1 + conv_state_layer1, ssm_state_layer1 = dynamic_context.mamba_states_cache(layer_number=1) + assert torch.all(conv_state_layer1 == 10.0) + assert torch.all(ssm_state_layer1 == 20.0) + + # Retrieve states using mamba_states_cache for global layer 3 + conv_state_layer3, ssm_state_layer3 = dynamic_context.mamba_states_cache(layer_number=3) + assert torch.all(conv_state_layer3 == 30.0) + assert torch.all(ssm_state_layer3 == 40.0) + @pytest.mark.internal def test_calculate_and_store_log_probs(self): self._setup_model_parallel_group(1, 1) @@ -665,7 +890,7 @@ def test_calculate_and_store_log_probs(self): num_attention_heads=2, max_sequence_length=512, buffer_size_gb=0.03, - buffer_guarenteed_fraction=0.1, + buffer_guaranteed_fraction=0.1, block_size_tokens=128, max_requests_override=None, max_tokens_override=None, diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine.py b/tests/unit_tests/inference/engines/test_dynamic_engine.py index 4ce8a5b11db..3ec7f80499b 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine.py @@ -4,7 +4,7 @@ import random import types from dataclasses import dataclass -from typing import Dict, List, Optional +from typing import Dict, List, Optional, Tuple import pytest import torch @@ -37,13 +37,29 @@ get_gpt_layer_with_transformer_engine_spec, ) from megatron.core.models.gpt.gpt_model import GPTModel +from megatron.core.models.mamba.mamba_layer_specs import mamba_stack_spec +from megatron.core.models.mamba.mamba_model import MambaModel from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.cuda_graphs import CudaGraphManager, _CudagraphGlobalRecord from megatron.core.transformer.transformer_config import TransformerConfig -from megatron.core.utils import is_fa_min_version +from megatron.core.utils import ( + check_mamba_sequence_packing_support, + get_attr_wrapped_model, + is_fa_min_version, + is_te_min_version, +) from tests.unit_tests.test_utilities import Utils +def skip_if_mamba_sequence_packing_not_available(model_provider: str): + if model_provider == "mamba": + sequence_packing_available, reason_for_no_sequence_packing = ( + check_mamba_sequence_packing_support() + ) + if not sequence_packing_available: + pytest.skip(reason_for_no_sequence_packing) + + def set_rounder(value): """Utility function to set the DynamicInferenceContext rounder.""" DynamicInferenceContext.ROUNDER = value # For backwards compatibility @@ -81,6 +97,9 @@ class DynamicEngineTestConfig: use_fixed_output_lengths: bool = False num_cuda_graphs: int = None + use_cuda_graphs_for_non_decode_steps: bool = True + fp8: bool = False + model_provider: str = "gpt" return_log_probs: bool = False materialize_only_last_token_logits: bool = True skip_prompt_log_probs_for_dynamic_inference: bool = False @@ -197,6 +216,9 @@ def _build_inference_context( test_config: DynamicEngineTestConfig, transformer_config: TransformerConfig, requests: List[DynamicInferenceRequest], + layer_type_list: Optional[List[str]], + mamba_conv_states_shape: Optional[Tuple[int]] = None, + mamba_ssm_states_shape: Optional[Tuple[int]] = None, ): """The inference context manages the KV cache and other inference state.""" @@ -208,6 +230,7 @@ def _build_inference_context( num_attention_heads=transformer_config.num_query_groups, max_sequence_length=test_config.max_sequence_length, num_cuda_graphs=test_config.num_cuda_graphs, + use_cuda_graphs_for_non_decode_steps=not test_config.model_provider == "mamba", buffer_size_gb=test_config.context_buffer_size_gb, buffer_guaranteed_fraction=test_config.context_buffer_guaranteed_fraction, block_size_tokens=test_config.context_block_size_tokens, @@ -215,6 +238,9 @@ def _build_inference_context( max_requests_override=test_config.context_max_requests_override, max_tokens_override=test_config.context_max_tokens_override, tensor_model_parallel_size=transformer_config.tensor_model_parallel_size, + layer_type_list=layer_type_list, + mamba_conv_states_shape=mamba_conv_states_shape, + mamba_ssm_states_shape=mamba_ssm_states_shape, materialize_only_last_token_logits=test_config.materialize_only_last_token_logits, use_flashinfer_fused_rope=None, # default to using flash-infer if available # this is for compatibility with the LTS environment @@ -223,6 +249,7 @@ def _build_inference_context( return context @classmethod + @torch.inference_mode() def _build_test_env(cls, test_config): Utils.initialize_model_parallel( tensor_model_parallel_size=test_config.tensor_model_parallel_size, @@ -241,72 +268,142 @@ def _build_test_env(cls, test_config): force_reset_rng=True, ) - # Transformer config. - transformer_config = TransformerConfig( - params_dtype=torch.bfloat16, - num_layers=4, - hidden_size=128 if test_config.fp8 else 32, - num_attention_heads=4, - use_cpu_initialization=True, - cuda_graph_impl=( - "local" - if test_config.num_cuda_graphs is not None and test_config.force_build_cuda_graphs - else "none" - ), - inference_rng_tracker=True, - tensor_model_parallel_size=test_config.tensor_model_parallel_size, - pipeline_model_parallel_size=test_config.pipeline_model_parallel_size, - expert_model_parallel_size=test_config.expert_model_parallel_size, - num_moe_experts=( - None - if test_config.expert_model_parallel_size == 1 - else test_config.expert_model_parallel_size - ), - sequence_parallel=test_config.sequence_parallel, - pipeline_dtype=torch.bfloat16, - add_bias_linear=test_config.expert_model_parallel_size == 1, - inference_sampling_seed=test_config.random_seed, - cuda_graph_scope=test_config.cuda_graph_scope, - ) - if test_config.fp8: - transformer_config.fp8 = "hybrid" - transformer_config.fp8_recipe = "tensorwise" - layer_spec = get_gpt_layer_with_transformer_engine_spec() - else: - layer_spec = get_gpt_layer_local_spec() - # Requests. requests = cls._build_requests(test_config) - # GPT model. - model = GPTModel( - config=transformer_config, - transformer_layer_spec=layer_spec, - vocab_size=test_config.vocab_size, - max_sequence_length=test_config.max_sequence_length, - parallel_output=True, - pre_process=parallel_state.is_pipeline_first_stage(), - post_process=parallel_state.is_pipeline_last_stage(), - ).cuda() + if test_config.model_provider == "gpt": + # Transformer config. + transformer_config = TransformerConfig( + params_dtype=torch.bfloat16, + num_layers=4, + hidden_size=128 if test_config.fp8 else 32, + num_attention_heads=4, + use_cpu_initialization=True, + cuda_graph_impl=( + "local" + if test_config.num_cuda_graphs is not None + and test_config.force_build_cuda_graphs + else "none" + ), + inference_rng_tracker=True, + tensor_model_parallel_size=test_config.tensor_model_parallel_size, + pipeline_model_parallel_size=test_config.pipeline_model_parallel_size, + expert_model_parallel_size=test_config.expert_model_parallel_size, + num_moe_experts=( + None + if test_config.expert_model_parallel_size == 1 + else test_config.expert_model_parallel_size + ), + sequence_parallel=test_config.sequence_parallel, + pipeline_dtype=torch.bfloat16, + add_bias_linear=test_config.expert_model_parallel_size == 1, + fp8="hybrid" if test_config.fp8 else None, + fp8_recipe="tensorwise" if test_config.fp8 else None, + inference_sampling_seed=test_config.random_seed, + cuda_graph_scope=test_config.cuda_graph_scope, + ) + if test_config.fp8: + layer_spec = get_gpt_layer_with_transformer_engine_spec() + else: + layer_spec = get_gpt_layer_local_spec() + + # GPT model. + model = GPTModel( + config=transformer_config, + transformer_layer_spec=layer_spec, + vocab_size=test_config.vocab_size, + max_sequence_length=test_config.max_sequence_length, + parallel_output=True, + pre_process=parallel_state.is_pipeline_first_stage(), + post_process=parallel_state.is_pipeline_last_stage(), + ).cuda() + elif test_config.model_provider == "mamba": + # Transformer config. + transformer_config = TransformerConfig( + params_dtype=torch.bfloat16, + num_layers=3, # 1 Mamba layer, 1 attention layer, 1 MLP layer + hidden_size=256, # The Mamba layer places several constraints on this + mamba_num_heads=16, + num_attention_heads=16, + use_cpu_initialization=True, + cuda_graph_impl=( + "local" + if test_config.num_cuda_graphs is not None + and test_config.force_build_cuda_graphs + else "none" + ), + inference_rng_tracker=True, + tensor_model_parallel_size=test_config.tensor_model_parallel_size, + pipeline_model_parallel_size=test_config.pipeline_model_parallel_size, + expert_model_parallel_size=test_config.expert_model_parallel_size, + num_moe_experts=( + None + if test_config.expert_model_parallel_size == 1 + else test_config.expert_model_parallel_size + ), + sequence_parallel=test_config.sequence_parallel, + pipeline_dtype=torch.bfloat16, + add_bias_linear=test_config.expert_model_parallel_size == 1, + fp8="hybrid" if test_config.fp8 else None, + fp8_recipe="tensorwise" if test_config.fp8 else None, + cuda_graph_scope=test_config.cuda_graph_scope, + ) + + # Mamba model. + model = MambaModel( + config=transformer_config, + mamba_stack_spec=mamba_stack_spec, + vocab_size=test_config.vocab_size, + max_sequence_length=test_config.max_sequence_length, + parallel_output=True, + hybrid_attention_ratio=0.3, + hybrid_mlp_ratio=0.3, + pre_process=parallel_state.is_pipeline_first_stage(), + post_process=parallel_state.is_pipeline_last_stage(), + ).cuda() + else: + raise ValueError(f"Invalid model provider {test_config.model_provider}") for param in model.parameters(): param.data = param.data.to(transformer_config.params_dtype) model.eval() + # Layer type list for hybrid models + decoder = get_attr_wrapped_model(model, "decoder") + layer_type_list = getattr(decoder, "layer_type_list", None) + if test_config.model_provider == "mamba": + mamba_states_shapes = decoder.mamba_state_shapes_per_request() + if mamba_states_shapes is not None: + (mamba_conv_states_shape, mamba_ssm_states_shape) = mamba_states_shapes + else: + # A `MambaBlock` can only not have a `MambaLayer` if using pipeline parallelism + # and a particular pipeline stage was not assigned a `MambaLayer`. + assert test_config.pipeline_model_parallel_size > 1 + mamba_conv_states_shape = None + mamba_ssm_states_shape = None + else: + mamba_conv_states_shape = None + mamba_ssm_states_shape = None + # Inference config. inference_config = InferenceWrapperConfig( hidden_size=transformer_config.hidden_size, inference_batch_times_seqlen_threshold=400, fp32_residual_connection=False, params_dtype=transformer_config.params_dtype, + fp8=transformer_config.fp8, padded_vocab_size=test_config.vocab_size, - fp8="hybrid" if test_config.fp8 else None, ) # Inference context. inference_context = cls._build_inference_context( - test_config=test_config, transformer_config=transformer_config, requests=requests + test_config=test_config, + transformer_config=transformer_config, + requests=requests, + layer_type_list=layer_type_list, + mamba_conv_states_shape=mamba_conv_states_shape, + mamba_ssm_states_shape=mamba_ssm_states_shape, ) # Inference model wrapper. @@ -348,6 +445,7 @@ def mock_detokenize_prompt(tokens): return env @classmethod + @torch.inference_mode() def _run_step(cls, env): set_rounder(4) # Step inference engine (i.e., generate one token per request). @@ -359,8 +457,8 @@ def _run_step(cls, env): finished_requests = result["finished_requests"] @classmethod + @torch.inference_mode() def _run_test(cls, **test_config_kwargs): - # Test environment. test_config = DynamicEngineTestConfig(**test_config_kwargs) env = cls._build_test_env(test_config) @@ -415,13 +513,16 @@ def teardown_method(self, method): @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) + @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) @pytest.mark.parametrize("num_cuda_graphs", [None, 1, 4]) @pytest.mark.parametrize("cuda_graph_scope", ["full", "full_iteration"]) - def test_simple(self, num_cuda_graphs, cuda_graph_scope) -> None: + def test_simple(self, model_provider, num_cuda_graphs, cuda_graph_scope) -> None: """Simple test that runs without errors, and validates output.""" + skip_if_mamba_sequence_packing_not_available(model_provider) # Run test. env = self._run_test( + model_provider=model_provider, num_cuda_graphs=num_cuda_graphs, context_max_requests_override=32, cuda_graph_scope=cuda_graph_scope, @@ -432,8 +533,8 @@ def test_simple(self, num_cuda_graphs, cuda_graph_scope) -> None: assert env.engine.context.max_requests == 32 assert env.engine.context.max_tokens == 160 - # Validate generated tokens. - expected_generated_tokens_list = [ + # Validate output tokens. + gpt_expected_generated_tokens = [ [69, 85, 55, 74], [29, 54, 85, 89], [33, 30, 64, 59], @@ -444,7 +545,26 @@ def test_simple(self, num_cuda_graphs, cuda_graph_scope) -> None: [], # this request is failed due to max sequence length overflow ] + mamba_expected_generated_tokens = [ + [74, 72, 83, 59], + [25, 54, 1, 70], + [28, 14, 15, 89], + [87, 27, 30, 52], + [44, 13, 82, 70], + [28, 74, 64, 16], + [8, 4, 83, 5], + [], + ] + + if model_provider == "gpt": + expected_generated_tokens_list = gpt_expected_generated_tokens + elif model_provider == "mamba": + expected_generated_tokens_list = mamba_expected_generated_tokens + else: + raise ValueError(f"Invalid model_provider {model_provider}") + assert len(env.requests) == len(expected_generated_tokens_list) + for request, expected_generated_tokens in zip(env.requests, expected_generated_tokens_list): assert request.generated_tokens == expected_generated_tokens, ( f"request {request.request_id}, " @@ -456,30 +576,41 @@ def test_simple(self, num_cuda_graphs, cuda_graph_scope) -> None: @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) - def test_overflow_factor(self) -> None: + def test_overflow_factor(self, model_provider: str = "gpt") -> None: """Test overflow factor arg.""" + skip_if_mamba_sequence_packing_not_available(model_provider) + # Run test. env = self._run_test( context_buffer_overflow_factor=0.1, context_max_requests_override=None, context_max_tokens_override=None, + model_provider=model_provider, ) # Validate max_requests, max_tokens. - assert env.engine.context.max_requests == 420 - assert env.engine.context.max_tokens == 420 + if model_provider == "gpt": + assert env.engine.context.max_requests == 420 + assert env.engine.context.max_tokens == 420 + elif model_provider == "mamba": + assert env.engine.context.max_requests == 16 + assert env.engine.context.max_tokens == 16 @pytest.mark.internal @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) - def test_request_overflow(self) -> None: + @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + def test_request_overflow(self, model_provider: str) -> None: """Test request overflow.""" - self._run_test(context_max_requests_override=4) + skip_if_mamba_sequence_packing_not_available(model_provider) + + self._run_test(context_max_requests_override=4, model_provider=model_provider) @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) + @torch.inference_mode() def test_token_overflow_transient(self) -> None: """Test token overflow.""" test_config = DynamicEngineTestConfig( @@ -515,13 +646,17 @@ def test_token_overflow_nontransient(self) -> None: @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) - def test_block_overflow(self) -> None: + @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + def test_block_overflow(self, model_provider: str) -> None: """Test block overflow.""" - env = self._build_test_env(DynamicEngineTestConfig()) + skip_if_mamba_sequence_packing_not_available(model_provider) + env = self._build_test_env(DynamicEngineTestConfig(model_provider=model_provider)) context = env.engine.context block_size_bytes = context.block_size_bytes buffer_size_gb = (block_size_bytes + 1) / 1024**3 - test_config = DynamicEngineTestConfig(context_buffer_size_gb=buffer_size_gb) + test_config = DynamicEngineTestConfig( + context_buffer_size_gb=buffer_size_gb, model_provider=model_provider + ) env = self._build_test_env(test_config) env.engine._add_request(env.requests[0]) assert list(env.engine.waiting_request_ids) == [0] @@ -530,17 +665,21 @@ def test_block_overflow(self) -> None: @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) - def test_multi_add(self) -> None: + @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + def test_multi_add(self, model_provider: str) -> None: """Test adding multiple requests simultaneously.""" - self._run_test(num_gap_steps=0) + skip_if_mamba_sequence_packing_not_available(model_provider) + self._run_test(num_gap_steps=0, model_provider=model_provider) @pytest.mark.internal @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) - def test_fixed_output_lengths(self) -> None: + @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + def test_fixed_output_lengths(self, model_provider: str) -> None: """Test generating a fixed number of output tokens.""" - self._run_test(use_fixed_output_lengths=True) + skip_if_mamba_sequence_packing_not_available(model_provider) + self._run_test(use_fixed_output_lengths=True, model_provider=model_provider) @pytest.mark.internal @pytest.mark.skipif( @@ -597,6 +736,7 @@ def test_cuda_graph_token_counts(self) -> None: (32, 32), ], ) + @torch.inference_mode() def test_cuda_graph_warmup( self, warmup_engine_mode: WarmupEngineMode, @@ -683,11 +823,18 @@ def test_cuda_graph_warmup( @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) - def test_generate_function(self) -> None: + @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + @torch.inference_mode() + def test_generate_function(self, model_provider: str) -> None: """Test the generate function that processes multiple prompts at once.""" + skip_if_mamba_sequence_packing_not_available(model_provider) + # Set up test environment test_config = DynamicEngineTestConfig( - num_requests=4, max_prompt_length=8, num_tokens_to_generate=4 + num_requests=4, + max_prompt_length=8, + num_tokens_to_generate=4, + model_provider=model_provider, ) env = self._build_test_env(test_config) @@ -733,36 +880,55 @@ async def test_run_engine(self): Test asynchronously adding and waiting for requests while the engine is running continuously. """ - # Test environment. - test_config = DynamicEngineTestConfig(use_fixed_output_lengths=True) - env = self._build_test_env(test_config) + # Have to wrap inference mode in-line because async functions are not supported + with torch.inference_mode(): + # Test environment. + test_config = DynamicEngineTestConfig(num_requests=8, use_fixed_output_lengths=True) + env = self._build_test_env(test_config) + + engine_task = asyncio.create_task(env.engine.run_engine(verbose=False)) + + request_completion_futures: Dict[int, asyncio.Future[DynamicInferenceRequest]] = {} + + # Add requests to engine. + for request in tqdm(env.requests, "add requests"): + request_completion_futures[request.request_id] = env.engine._add_request(request) + + # Wait for all requests to complete. + await asyncio.gather(*request_completion_futures.values()) + + # Verify that all request outputs were set. + for request_id, fut in request_completion_futures.items(): + num_tokens_to_generate = env.requests[ + request_id + ].sampling_params.num_tokens_to_generate + result = fut.result() + assert result.generated_length == num_tokens_to_generate, ( + f"Request {request_id} expected to generate {num_tokens_to_generate} " + f"tokens but generated {result.generated_length}" + ) - engine_task = asyncio.create_task(env.engine.run_engine(verbose=False)) + engine_task.cancel() - request_completion_futures: Dict[int, asyncio.Future[DynamicInferenceRequest]] = {} + @pytest.mark.internal + @pytest.mark.skipif( + not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" + ) + @pytest.mark.skipif(not is_te_min_version("2.2.0"), reason="TE 2.2.0 is required") + @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + def test_fp8_inference(self, model_provider: str): + skip_if_mamba_sequence_packing_not_available(model_provider) - # Add requests to engine. - for request in tqdm(env.requests, "add requests"): - request_completion_futures[request.request_id] = env.engine._add_request(request) - - # Wait for all requests to complete. - await asyncio.gather(*request_completion_futures.values()) - - # Verify that all request outputs were set. - for request_id, fut in request_completion_futures.items(): - num_tokens_to_generate = env.requests[request_id].sampling_params.num_tokens_to_generate - result = fut.result() - assert result.generated_length == num_tokens_to_generate, ( - f"Request {request_id} expected to generate {num_tokens_to_generate} " - f"tokens but generated {result.generated_length}" - ) + fp8_available, reason_for_no_fp8 = check_fp8_support() + if not fp8_available: + pytest.skip(reason_for_no_fp8) - engine_task.cancel() + self._run_test(model_provider=model_provider, fp8=True) - @pytest.mark.internal @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) + @torch.inference_mode() def test_return_log_probs(self): """Verify that returning log probs does not raise any error.""" # Returning log probs requires materializing the full prompt logits or @@ -785,9 +951,19 @@ def test_return_log_probs(self): @pytest.mark.parametrize("ep_size", [1, 2]) @pytest.mark.parametrize("pp_size", [1, 2]) @pytest.mark.parametrize("tp_size", [1, 2]) + @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + @torch.inference_mode() def test_parallel_inference( - self, tp_size, pp_size, ep_size, sequence_parallel, materialize_only_last_token_logits + self, + model_provider, + tp_size, + pp_size, + ep_size, + sequence_parallel, + materialize_only_last_token_logits, ): + skip_if_mamba_sequence_packing_not_available(model_provider) + if tp_size == 1 and pp_size == 1 and ep_size == 1: pytest.skip(reason="Test requires tp_size > 1 or pp_size > 1 or ep_size > 1") elif not torch.distributed.is_initialized(): @@ -800,7 +976,16 @@ def test_parallel_inference( pytest.skip(reason="Sequence parallelism requires tp_size > 1") elif tp_size > 1 and ep_size > 1 and not sequence_parallel: pytest.skip(reason="Sequence parallelism must be used with tp_size > 1 and ep_size > 1") + elif pp_size > 1 and model_provider == "mamba": + pytest.skip( + reason=( + "Running hybrid models with pp_size > 1 and no attention on some " + "pipeline stages is not supported yet." + ) + ) + env = self._run_test( + model_provider=model_provider, tensor_model_parallel_size=tp_size, pipeline_model_parallel_size=pp_size, expert_model_parallel_size=ep_size, @@ -881,6 +1066,32 @@ def test_events(self): assert result_event_types == expected_event_types + @pytest.mark.internal + @pytest.mark.skipif( + not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" + ) + @pytest.mark.parametrize("model_provider", ["gpt", "mamba"]) + @torch.inference_mode() + def test_chunked_prefill(self, model_provider: str): + """Verify that chunked prefill output is equivalent to regular prefill.""" + skip_if_mamba_sequence_packing_not_available(model_provider) + + prompt_length = 1200 + num_tokens_to_generate = 16 + max_sequence_length = prompt_length + num_tokens_to_generate + + # Configure context to force chunking (chunked prefill is enabled by default) + env = self._run_test( + num_requests=1, + min_prompt_length=prompt_length, + max_prompt_length=prompt_length, + num_tokens_to_generate=num_tokens_to_generate, + materialize_only_last_token_logits=False, + model_provider=model_provider, + context_block_size_tokens=256, + context_max_tokens_override=300, + ) + if __name__ == "__main__": test = TestDynamicInferenceEngine() diff --git a/tools/run_inference_performance_test.py b/tools/run_inference_performance_test.py index 2f2adabc0ab..01e5ab58898 100644 --- a/tools/run_inference_performance_test.py +++ b/tools/run_inference_performance_test.py @@ -1,83 +1,59 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -import os -from megatron.core.inference.model_inference_wrappers.inference_wrapper_config import ( - InferenceWrapperConfig, -) import argparse +import os import random -import torch import sys import time -import tqdm -import warnings -from model_provider import model_provider + +import torch + from gpt_builders import gpt_builder from mamba_builders import mamba_builder -from megatron.core.inference.engines.abstract_engine import AbstractEngine +from megatron.core.inference.contexts import DynamicInferenceContext from megatron.core.inference.engines import DynamicInferenceEngine, StaticInferenceEngine +from megatron.core.inference.engines.abstract_engine import AbstractEngine from megatron.core.inference.inference_request import InferenceRequest -from megatron.core.inference.contexts import DynamicInferenceContext -from megatron.core.inference.sampling_params import SamplingParams from megatron.core.inference.model_inference_wrappers.gpt.gpt_inference_wrapper import ( GPTInferenceWrapper, ) +from megatron.core.inference.model_inference_wrappers.inference_wrapper_config import ( + InferenceWrapperConfig, +) +from megatron.core.inference.sampling_params import SamplingParams from megatron.core.inference.text_generation_controllers.text_generation_controller import ( TextGenerationController, ) +from megatron.core.ssm.mamba_hybrid_layer_allocation import Symbols from megatron.core.transformer.module import MegatronModule +from megatron.core.utils import get_attr_wrapped_model +from model_provider import model_provider sys.path.append( os.path.abspath(os.path.join(os.path.dirname(__file__), os.path.pardir, os.path.pardir)) ) -from megatron.training import get_args -from megatron.training import get_tokenizer -from megatron.training.checkpointing import load_checkpoint -from megatron.core import mpu -from megatron.training.initialize import initialize_megatron -from megatron.training import get_model, get_tokenizer import asyncio from functools import partial -from typing import AsyncIterator, List, Union +from typing import List, Union + +from examples.inference.gpt.utils import add_common_inference_args +from megatron.core import mpu +from megatron.training import get_args, get_model, get_tokenizer +from megatron.training.checkpointing import load_checkpoint +from megatron.training.initialize import initialize_megatron REQUEST_ID = 0 -def add_text_generate_args(parser): - """Text generation arguments.""" - group = parser.add_argument_group(title='text generation') +def add_inference_benchmarking_args(parser): + """Inference benchmarking arguments.""" + parser = add_common_inference_args(parser) + + group = parser.add_argument_group(title='inference_benchmarking') - group.add_argument("--temperature", type=float, default=1.0, help='Sampling temperature.') - group.add_argument("--top_k", type=int, default=1, help='Top k sampling.') - group.add_argument("--top_p", type=float, default=0.0, help='Top p sampling.') - group.add_argument( - "--return-log-probs", - action='store_true', - default=False, - help='Return the log probabilities of the final output tokens', - ) - group.add_argument("--top-n-logprobs", type=int, default=0, help="Top-N logprobs") - group.add_argument( - "--num-tokens-to-generate", - type=int, - default=30, - help='Number of tokens to generate for each prompt', - ) - group.add_argument( - "--prompts", - metavar='N', - type=str, - default=None, - nargs='+', - help='Input prompts with each prompt within quotes and seperated by space', - ) - group.add_argument( - "--num-input-tokens", type=int, default=None, help='Number of input tokens per prompt' - ) - group.add_argument("--stream", action="store_true", default=False, help="Stream output tokens") group.add_argument( - "--model-provider", choices=["mamba", "gpt"], default="gpt", help="Model provider" + "--num-input-tokens", type=int, default=128, help="Number of input tokens per request" ) group.add_argument( "--engine-type", choices=["static", "dynamic"], default="static", help="Engine type" @@ -85,14 +61,13 @@ def add_text_generate_args(parser): group.add_argument( "--benchmark-profile", action="store_true", default=False, help="If set, profile" ) + group.add_argument('--stream', action="store_true", default=False, help="If set, stream tokens") return parser def get_inference_engine(args: argparse.Namespace, model: MegatronModule) -> AbstractEngine: """Utility to get the relevant backend for running inference - This function will automatically chose the TRTLLMBackend when possible, and if not revert to Mcore backend if the user does not specify any backends. TRT LLM Backend is not implmented yet. - Args: args (Namespace): The user arguments parsed from command line model (MegatronModule): The megatron model . @@ -111,9 +86,18 @@ def get_inference_engine(args: argparse.Namespace, model: MegatronModule) -> Abs inference_max_requests=args.inference_max_batch_size, inference_max_seq_length=args.inference_max_seq_length, nccl_all_reduce_for_prefill=args.nccl_all_reduce_for_prefill, - moe_pad_experts_for_cuda_graph_inference = args.moe_pad_experts_for_cuda_graph_inference + moe_pad_experts_for_cuda_graph_inference=args.moe_pad_experts_for_cuda_graph_inference, ) + # Layer type list for hybrid models + decoder = get_attr_wrapped_model(model, "decoder") + layer_type_list = getattr(decoder, "layer_type_list", None) + if layer_type_list is not None and Symbols.MAMBA in layer_type_list: + (mamba_conv_states_shape, mamba_ssm_states_shape) = decoder.mamba_state_shapes_per_request() + else: + mamba_conv_states_shape = None + mamba_ssm_states_shape = None + if args.engine_type == "static": inference_wrapped_model = GPTInferenceWrapper(model, inference_wrapper_config) inference_wrapped_model.model_is_pipeline_parallel = not ( @@ -132,12 +116,28 @@ def get_inference_engine(args: argparse.Namespace, model: MegatronModule) -> Abs args.num_query_groups if args.group_query_attention else args.num_attention_heads ), max_sequence_length=args.inference_max_seq_length, + num_cuda_graphs=( + args.inference_dynamic_batching_num_cuda_graphs + if args.cuda_graph_impl == "local" + else None + ), buffer_size_gb=args.inference_dynamic_batching_buffer_size_gb, buffer_guaranteed_fraction=args.inference_dynamic_batching_buffer_guaranteed_fraction, buffer_overflow_factor=args.inference_dynamic_batching_buffer_overflow_factor, max_requests_override=args.inference_dynamic_batching_max_requests_override, max_tokens_override=args.inference_dynamic_batching_max_tokens_override, block_size_tokens=args.inference_dynamic_batching_block_size, + tensor_model_parallel_size=args.tensor_model_parallel_size, + materialize_only_last_token_logits=not args.return_log_probs, + layer_type_list=layer_type_list, + mamba_conv_states_shape=mamba_conv_states_shape, + mamba_ssm_states_shape=mamba_ssm_states_shape, + cache_mla_latent=args.multi_latent_attention and args.cache_mla_latents, + kv_lora_rank=args.kv_lora_rank if args.multi_latent_attention else None, + qk_pos_emb_head_dim=args.qk_pos_emb_head_dim, + use_cuda_graphs_for_non_decode_steps=not args.decode_only_cuda_graphs, + use_flashinfer_fused_rope=args.use_flashinfer_fused_rope, + unified_memory_level=args.inference_dynamic_batching_unified_memory_level, ) inference_wrapped_model = GPTInferenceWrapper( model, inference_wrapper_config, inference_context=context @@ -269,7 +269,7 @@ def main(): # Note: The default args passed here can be overwritten by using appropriate params (check arguments.py file) # Micro batch size is not needed to be set by user. (It is calculated based on inference-batch-times-seqlen-threshold argument) initialize_megatron( - extra_args_provider=add_text_generate_args, + extra_args_provider=add_inference_benchmarking_args, args_defaults={ 'no_load_rng': True, 'no_load_optim': True, @@ -285,6 +285,8 @@ def main(): model_builder = gpt_builder elif args.model_provider == "mamba": model_builder = mamba_builder + else: + raise ValueError(f"Invalid model provider {args.model_provider}") model = get_model(partial(model_provider, model_builder), wrap_with_ddp=False) tokenizer = get_tokenizer() @@ -338,10 +340,7 @@ def main(): print(f"Running warmup for CUDA graphs...") warmup_sampling_params = SamplingParams(num_tokens_to_generate=10) warmup_sampling_params.add_attributes({"no_early_termination": True}) - if args.engine_type == "static": - inference_engine.generate(prompts=["warmup"], sampling_params=warmup_sampling_params) - elif args.engine_type == "dynamic": - generate_dynamic(args, requests, inference_engine) + inference_engine.generate(prompts=["warmup"], sampling_params=warmup_sampling_params) if args.benchmark_profile: torch.cuda.cudart().cudaProfilerStart()