diff --git a/megatron/core/inference/contexts/dynamic_context.py b/megatron/core/inference/contexts/dynamic_context.py index 2f559bf581d..10bef5b3c15 100644 --- a/megatron/core/inference/contexts/dynamic_context.py +++ b/megatron/core/inference/contexts/dynamic_context.py @@ -238,9 +238,9 @@ class DynamicInferenceContext(BaseInferenceContext): use_flashinfer_fused_rope (bool): If True, use flashinfer's fused rope implementation. If None, defaults to using flash-infer if available. metrics_writer (Optional['WandbModule']): Wandb module for writing metrics. - num_request_metadata (Optional[int]): Number of metadata fields to track per request. - These represent metadata that is needed by the text generation controller, - and that must be kept in sync with active requests through update_requests. + request_metadata_types (Optional[List[Tuple[str, torch.dtype, bool]]]): A list of the + per-request metadata types to track. Each entry is a tuple consisting of the string + label, the target dtype, and whether to store the data on GPU. """ DEFAULT_MAX_TOKENS = 16384 @@ -271,7 +271,7 @@ def __init__( cuda_graph_max_tokens: Optional[int] = None, cuda_graph_mixed_prefill_count: Optional[int] = 16, metrics_writer: Optional['WandbModule'] = None, - num_request_metadata: Optional[int] = None, + request_metadata_types: Optional[List[Tuple[str, torch.dtype, bool]]] = None, ): super().__init__(materialize_only_last_token_logits=materialize_only_last_token_logits) @@ -399,9 +399,9 @@ def __init__( ) # Track request metadata. - if num_request_metadata is None: - num_request_metadata = len(DynamicInferenceRequest.get_metadata_labels()) - self.num_request_metadata = num_request_metadata + if request_metadata_types is None: + request_metadata_types = DynamicInferenceRequest.get_metadata_types() + self.request_metadata_types = request_metadata_types # Initialize context state. self.params_dtype = params_dtype @@ -537,11 +537,12 @@ def allocate_all_tensors(self, *, is_init: bool) -> None: ) # Track request metadata. - self.request_metadata = torch.empty( - (self.max_total_requests, self.num_request_metadata), - dtype=torch.float32, - device=torch.cuda.current_device(), - ) + self.request_metadata = { + label: torch.empty( + (self.max_total_requests,), dtype=dtype, device=torch.cuda.current_device() + ) + for label, dtype, _ in self.request_metadata_types + } # Per-token state. self.token_to_input_ids = torch.full( @@ -999,7 +1000,7 @@ def add_dummy_requests_parallel( num_tokens_to_generate: List[int] = [] request_ids: List[int] = [] prompt_tokens: List[Tensor] = [] - metadata_rows: List[List[float]] = [] + metadata_cols: List[List] = [[] for _ in self.request_metadata_types] for req in requests: assert isinstance( @@ -1020,7 +1021,8 @@ def add_dummy_requests_parallel( device=self.token_to_input_ids.device, dtype=self.token_to_input_ids.dtype ) ) - metadata_rows.append(req.tracked_metadata) + for i, m in enumerate(req.tracked_metadata): + metadata_cols[i].append(m) total_new_tokens = sum(lengths) if self.active_token_count + total_new_tokens > self.max_tokens: @@ -1034,9 +1036,6 @@ def add_dummy_requests_parallel( num_tokens_to_generate, dtype=self.request_query_lengths.dtype, device=device ) request_ids_tensor = torch.tensor(request_ids, dtype=self.request_ids.dtype, device=device) - metadata_tensor = torch.tensor( - metadata_rows, dtype=self.request_metadata.dtype, device=self.request_metadata.device - ) block_counts = torch.div( lengths_tensor + (self.block_size_tokens - 1), @@ -1053,7 +1052,10 @@ def add_dummy_requests_parallel( self.request_output_lengths[request_slice] = lengths_tensor + tokens_to_generate_tensor self.request_kv_length_offsets[request_slice] = 0 self.request_kv_block_counts[request_slice] = block_counts - self.request_metadata[request_slice] = metadata_tensor + for i, (label, dtype, _) in enumerate(self.request_metadata_types): + self.request_metadata[label][request_slice] = torch.tensor( + metadata_cols[i], dtype=dtype, device=torch.cuda.current_device() + ) dummy_block_idx = self.block_allocator.dummy_block_idx self.request_last_kv_block_id[request_slice] = dummy_block_idx @@ -1325,7 +1327,10 @@ def reset(self) -> None: self.request_last_kv_block_id.fill_(-1) self.request_last_kv_block_offset.fill_(0) self.request_to_kv_block_ids.fill_(-1) - self.request_metadata.fill_(0) + + # Reset request metadata. + for metadata_tensor in self.request_metadata.values(): + metadata_tensor.fill_(0) # Reset token indexes. self.token_to_input_ids.fill_(0) @@ -1482,14 +1487,17 @@ def add_request(self, req: DynamicInferenceRequest, chunk_length: Optional[int] raise TokenOverflowError(req.request_id) self.request_ids[current_id] = req.request_id + # Handle request metadata. - metadata = req.tracked_metadata assert ( - len(metadata) == self.num_request_metadata - ), "Request added to context with invalid metadata length" - self.request_metadata[current_id] = torch.tensor( - metadata, dtype=torch.float32, device=self.request_metadata.device - ) + req.get_metadata_types() == self.request_metadata_types + ), "Request added to context with invalid metadata types" + metadata = req.tracked_metadata + metadata_types = req.get_metadata_types() + for m, m_type in zip(metadata, metadata_types): + label, _, _ = m_type + self.request_metadata[label][current_id] = m + # Handle length and block assignments. self.request_query_lengths[current_id] = chunk_length self.request_output_lengths[current_id] = ( @@ -1556,7 +1564,6 @@ def _move_book_keeping_tensors(self, src_idxs, dst_idxs, next_tokens): self.request_kv_length_offsets[dst_idxs] = self.request_kv_length_offsets[src_idxs] self.request_query_lengths[dst_idxs] = self.request_query_lengths[src_idxs] self.request_output_lengths[dst_idxs] = self.request_output_lengths[src_idxs] - self.request_metadata[dst_idxs] = self.request_metadata[src_idxs] self.request_ids[dst_idxs] = self.request_ids[src_idxs] next_tokens[dst_idxs] = next_tokens[src_idxs] @@ -1565,6 +1572,9 @@ 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] + for metadata_tensor in self.request_metadata.values(): + metadata_tensor[dst_idxs] = metadata_tensor[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] @@ -1577,7 +1587,6 @@ def _swap_book_keeping_tensors(self, src_idxs, dst_idxs, next_tokens): tensor_swap(self.request_kv_length_offsets, src_idxs, dst_idxs) tensor_swap(self.request_query_lengths, src_idxs, dst_idxs) tensor_swap(self.request_output_lengths, src_idxs, dst_idxs) - tensor_swap(self.request_metadata, src_idxs, dst_idxs) tensor_swap(self.request_ids, src_idxs, dst_idxs) tensor_swap(next_tokens, src_idxs, dst_idxs) tensor_swap(self.request_to_kv_block_ids, src_idxs, dst_idxs) @@ -1585,6 +1594,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) + for metadata_tensor in self.request_metadata.values(): + tensor_swap(metadata_tensor, src_idxs, dst_idxs) + if self.is_hybrid_model: tensor_swap(self.mamba_metadata.request_to_mamba_state_idx, src_idxs, dst_idxs) diff --git a/megatron/core/inference/inference_request.py b/megatron/core/inference/inference_request.py index 6d0ff898bad..458fbad387f 100644 --- a/megatron/core/inference/inference_request.py +++ b/megatron/core/inference/inference_request.py @@ -6,12 +6,13 @@ import warnings from dataclasses import asdict, dataclass, field from enum import Enum, auto -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Optional, Tuple import torch from megatron.core.inference.sampling_params import SamplingParams from megatron.core.tokenizers import MegatronTokenizer +from megatron.core.utils import experimental_api def serialize_tensor(tensor: torch.Tensor) -> bytes: @@ -228,6 +229,7 @@ def deserialize(cls, obj: dict) -> "DynamicInferenceEvent": return event +@experimental_api @dataclass(kw_only=True) class DynamicInferenceRequest(InferenceRequest): """Class for one inference request @@ -313,21 +315,27 @@ def tracked_metadata(self) -> List[Any]: "in its sampling_params. Defaulting to -1." ) sp.termination_id = -1 - return [getattr(sp, field) for field in self.get_metadata_labels().keys()] + return [getattr(sp, field) for field, _, _ in self.get_metadata_types()] @staticmethod - def get_metadata_labels() -> Dict[str, int]: - """Provides human-readable labels for the tracked metadata fields.""" - ret = [ - "temperature", - "top_k", - "top_p", - "termination_id", - "return_log_probs", - "skip_prompt_log_probs", - "top_n_logprobs", + def get_metadata_types() -> List[Tuple[str, torch.dtype, bool]]: + """Keeps track of all request metadata names, dtypes, and target device. + + Returns: + List[Tuple[str, torch.dtype, bool]]: Mapping from metadata name to: + name (str) - The name of the metadata field. + dtype (torch.dtype) - The datatype of the metadata. + on_device (bool) - Whether the metadata lives on GPU (True) or CPU (False). + """ + return [ + ("temperature", torch.float32, False), # CPU for torch sampling + ("top_k", torch.int32, False), # CPU for torch sampling + ("top_p", torch.float32, False), # CPU for torch sampling + ("termination_id", torch.int64, True), + ("return_log_probs", torch.bool, False), # CPU for non-selective logprobs + ("skip_prompt_log_probs", torch.bool, False), # CPU for non-selective logprobs + ("top_n_logprobs", torch.int32, False), # CPU for torch sampling ] - return {k: v for v, k in enumerate(ret)} def add_event(self, type: DynamicInferenceEventType, payload: Optional[Any] = None) -> None: """Add event.""" diff --git a/megatron/core/inference/text_generation_controllers/text_generation_controller.py b/megatron/core/inference/text_generation_controllers/text_generation_controller.py index dcb8b419e74..2ff0dcd579e 100644 --- a/megatron/core/inference/text_generation_controllers/text_generation_controller.py +++ b/megatron/core/inference/text_generation_controllers/text_generation_controller.py @@ -6,7 +6,7 @@ import functools import inspect from collections import defaultdict -from typing import Any, Dict, List, Optional, OrderedDict, Tuple, Union +from typing import Any, Dict, Iterator, List, Optional, OrderedDict, Tuple, Union import torch import torch.nn.functional as F @@ -20,11 +20,7 @@ is_pipeline_last_stage, ) from megatron.core.inference.contexts.dynamic_context import MaxSequenceLengthOverflowError -from megatron.core.inference.inference_request import ( - DynamicInferenceRequest, - InferenceRequest, - Status, -) +from megatron.core.inference.inference_request import InferenceRequest, Status from megatron.core.inference.model_inference_wrappers.abstract_model_inference_wrapper import ( AbstractModelInferenceWrapper, ) @@ -91,22 +87,22 @@ def _init_dynamic_sampling_tensors(self): # Use padded vocab size because tokenizer vocab size might pad to nearest power of 2. vocab_size = self.inference_wrapped_model.inference_wrapper_config.padded_vocab_size - # Initialize bookkeeping tensors. - self.sampling_logits_cuda = torch.empty( - max_requests, vocab_size, dtype=logits_dtype, device=device - ) - self.sampled_tokens_cuda = torch.empty(max_requests, dtype=torch.int64, device=device) + self._sampling_backend = "torch" + self._sampled_tokens_cuda = torch.empty(max_requests, dtype=torch.int64, device=device) - self.temperature_cuda = torch.empty_like(self.sampled_tokens_cuda, dtype=torch.float) - self.top_k_cuda = torch.empty_like(self.sampled_tokens_cuda, dtype=torch.int32) - self.top_p_cuda = torch.empty_like(self.sampled_tokens_cuda, dtype=torch.float) - self.termination_id_cuda = torch.empty(max_requests, dtype=torch.int64, device=device) - self.return_log_probs_cuda = torch.empty(max_requests, dtype=torch.bool, device=device) - self.skip_prompt_log_probs_cuda = torch.empty(max_requests, dtype=torch.bool, device=device) - self.top_n_logprobs_cuda = torch.empty(max_requests, dtype=torch.int32, device=device) + # Keep track of request metadata. + self._request_metadata: Dict[str, Tensor] = {} + for label, dtype, on_gpu in context.request_metadata_types: + tensor = context.request_metadata[label] + if not on_gpu: + # Create pinned tensors for request metadata that lives on CPU. + # This is metadata which requires D2H copies, such as top_k for torch sampling. + tensor = torch.empty_like(tensor, device="cpu", pin_memory=True) + self._request_metadata[label] = tensor # Used for inefficient torch sampling. - self.torch_sampling_buckets: List[Tensor] = [] + if self._sampling_backend == "torch": + self._torch_sampling_buckets: Iterator[Tuple] = [] def tokenize_prompt(self, prompt: str, add_BOS: bool = False) -> List[int]: """Utility to tokenize the input prompts. @@ -492,6 +488,7 @@ def _dynamic_step_context_init( """ context = self.inference_wrapped_model.inference_context inference_wrapper_config = self.inference_wrapped_model.inference_wrapper_config + active_request_slice = slice(context.paused_request_count, context.total_request_count) # Remove Float16Module wrapper if it exists unwrapped_model = unwrap_model(self.inference_wrapped_model.model) @@ -523,6 +520,14 @@ def _dynamic_step_context_init( # Turn off symmetric all reduces for prefill unwrapped_model.set_symmetric_ar(None) + # Get request metadata for this step. + for label, dtype, on_gpu in context.request_metadata_types: + if not on_gpu: + # We need a D2H copy from the context to the pinned memory buffer. + self._request_metadata[label].copy_( + context.request_metadata[label], non_blocking=True + ) + # Get flat tokens, position ids. if construct_graph_dimensions is not None: return context.current_input_and_position_ids( @@ -543,8 +548,6 @@ def _dynamic_step_forward_logits(self, input_ids: Tensor, position_ids: Tensor) inference_wrapper_config = self.inference_wrapped_model.inference_wrapper_config context = self.inference_wrapped_model.inference_context - materialize_only_last_token_logits = context.materialize_only_last_token_logits - active_request_count = context.total_request_count - context.paused_request_count with torch.inference_mode(): @@ -554,7 +557,9 @@ def _dynamic_step_forward_logits(self, input_ids: Tensor, position_ids: Tensor) if self.model_is_pipeline_parallel: logits_seq_len = ( - active_request_count if materialize_only_last_token_logits else input_ids.shape[1] + active_request_count + if context.materialize_only_last_token_logits + else input_ids.shape[1] ) vocab_size = inference_wrapper_config.padded_vocab_size logits_shape = [1, logits_seq_len, vocab_size] @@ -568,178 +573,103 @@ def _dynamic_step_forward_logits(self, input_ids: Tensor, position_ids: Tensor) tensor=logits, pp_group=self.pp_group, ) - return logits - def _dynamic_step_sample_bookkeeping( - self, - *, - backend: str = "torch", - request_metadata: Optional[Tensor] = None, - request_metadata_labels: Dict[str, int] = None, - ): - """Perform bookkeeping necessary to sample logits for dynamic batching. - - The ability to override the context's data is solely intended for - standalone use or testing, and should never be used in a running system. + return logits - Args: - backend (str): The sampling backend to use. - request_metadata (Optional[Tensor]): An override for the tensor that manages all - request metadata, such as sampling parameters. By default, this metadata is - retrieved from the context. - request_metadata_labels (Optional[Dict]): An override for the map of metadata labels - to their index in the request_metadata tensor. By default, this metadata is - retrieved from the request object. - """ - assert backend in ["torch"] + def _dynamic_step_sample_bookkeeping(self): + """Perform bookkeeping necessary to sample logits for dynamic batching.""" context = self.inference_wrapped_model.inference_context + active_request_slice = slice(context.paused_request_count, context.total_request_count) - if request_metadata is None: - request_metadata = context.request_metadata[ - context.paused_request_count : context.total_request_count, : - ] - if request_metadata_labels is None: - request_metadata_labels = DynamicInferenceRequest.get_metadata_labels() - active_request_count = request_metadata.size(0) - - # Shorthand these, because the torch backend needs them. - temp = request_metadata[:, request_metadata_labels["temperature"]] - top_k = request_metadata[:, request_metadata_labels["top_k"]] - top_p = request_metadata[:, request_metadata_labels["top_p"]] - - # Copy data into relevant tensors. - self.temperature_cuda[:active_request_count].copy_(temp, non_blocking=True) - self.top_k_cuda[:active_request_count] = top_k.to( - dtype=torch.int32, copy=True, non_blocking=True - ) - self.top_p_cuda[:active_request_count].copy_(top_p, non_blocking=True) - self.termination_id_cuda[:active_request_count] = request_metadata[ - :, request_metadata_labels["termination_id"] - ].to(dtype=torch.int64, copy=True, non_blocking=True) - self.return_log_probs_cuda[:active_request_count] = request_metadata[ - :, request_metadata_labels["return_log_probs"] - ].to(dtype=torch.bool, copy=True, non_blocking=True) - self.skip_prompt_log_probs_cuda[:active_request_count] = request_metadata[ - :, request_metadata_labels["skip_prompt_log_probs"] - ].to(dtype=torch.bool, copy=True, non_blocking=True) - self.top_n_logprobs_cuda[:active_request_count] = request_metadata[ - :, request_metadata_labels["top_n_logprobs"] - ].to(dtype=torch.int32, copy=True, non_blocking=True) - - if backend == "torch": + if self._sampling_backend == "torch": # Bucketize the core sampling parameters. - core_params = torch.stack((temp, top_k, top_p), dim=1) - _, inv_indices, cnts = torch.unique( - core_params, dim=0, return_inverse=True, return_counts=True - ) - order = torch.argsort(inv_indices, stable=True) - sampling_buckets = torch.split(order, cnts.tolist()) - # Perform the D2H sync needed by `_torch_sampling_func` here. - group_reps = torch.stack([indices[0] for indices in sampling_buckets], dim=0) - core_params_reps = core_params[group_reps].detach().cpu() - temp_reps = core_params_reps[:, 0].tolist() - top_k_reps = core_params_reps[:, 1].to(torch.int32).tolist() - top_p_reps = core_params_reps[:, 2].tolist() + # Doing so via list comprehension is orders of magnitude faster than via torch. + bucket_map = {} + + # Shorthands for the dictionary comprehension. + temp = self._request_metadata["temperature"][active_request_slice].tolist() + top_k = self._request_metadata["top_k"][active_request_slice].tolist() + top_p = self._request_metadata["top_p"][active_request_slice].tolist() + + for i, (t, k, p) in enumerate(zip(temp, top_k, top_p)): + h = (t, k, p) + bucket = bucket_map.get(h, None) + if bucket is None: + bucket_map[h] = ([i], i) + else: + bucket[0].append(i) + # Store the buckets and their equivalence class representatives. - self.torch_sampling_buckets = ( - (sampling_buckets[idx], temp_reps[idx], top_k_reps[idx], top_p_reps[idx]) - for idx in range(len(sampling_buckets)) + self._torch_sampling_buckets = ( + (indices, temp[rep], top_k[rep], top_p[rep]) for indices, rep in bucket_map.values() ) - def _dynamic_step_sample_logits(self, logits: Tensor, backend: str = "torch") -> Tensor: + def _dynamic_step_sample_logits(self, logits: Tensor): """Sample tokens from logits for dynamic batching. Args: - logits (Tensor): The logits to sample from. - backend (str): The sampling backend to use. - - Returns: - new_sample (Tensor): The sampled tokens. + logits (Tensor): The logits from the forward pass. """ # TODO(ksanthanam): Evaluate whether it makes more sense to sample on 1 rank # and then broadcast the sampled tokens rather than broadcasting the raw logits. - assert backend in ["torch"] - - context = self.inference_wrapped_model.inference_context - materialize_only_last_token_logits = context.materialize_only_last_token_logits # Last token logits. - if materialize_only_last_token_logits: + context = self.inference_wrapped_model.inference_context + if context.materialize_only_last_token_logits: # When materialize_only_last_token_logits is true, last_token_logits is # already called in the forward pass of GPT. last_token_logits = logits.squeeze(0) else: last_token_logits = context.last_token_logits(logits) - active_request_count = last_token_logits.size(0) - # Copy last_token_logits to contiguous buffer. - self.sampling_logits_cuda[:active_request_count].copy_(last_token_logits, non_blocking=True) - if backend == "torch": + if self._sampling_backend == "torch": # Concatenate the outputs once to prevent repeated small writes. token_list = [] indices_list = [] - for indices, temp, top_k, top_p in self.torch_sampling_buckets: + for indices, temp, top_k, top_p in self._torch_sampling_buckets: token_list.append( - self._torch_sampling_func( - self.sampling_logits_cuda[indices, :], temp, top_k, top_p - ) + self._torch_sampling_func(last_token_logits[indices, :], temp, top_k, top_p) ) - indices_list.append(indices) + indices_list.append(torch.tensor(indices)) # Single write to the output tensor. sampled_tokens = torch.cat(token_list, dim=0) sampled_indices = torch.cat(indices_list, dim=0) - self.sampled_tokens_cuda.index_copy_(0, sampled_indices, sampled_tokens) - return self.sampled_tokens_cuda[:active_request_count].clone() - - def _dynamic_step_log_probs_bookkeeping(self) -> bool: - """Perform bookkeeping necessary to compute log probs for dynamic batching.""" - context = self.inference_wrapped_model.inference_context - materialize_only_last_token_logits = context.materialize_only_last_token_logits - - active_request_count = context.total_request_count - context.paused_request_count - - # Create a copy to avoid modifying the original tensor with in-place operations - to_check = self.return_log_probs_cuda[:active_request_count].clone() - to_check &= ~self.skip_prompt_log_probs_cuda[:active_request_count] + self._sampled_tokens_cuda[sampled_indices] = sampled_tokens - assert not ( - to_check.any() and materialize_only_last_token_logits - ), "Prompt log probs cannot be calculated if only last token logits are materialized. Set materialize_only_last_token_logits to False in DynamicInferenceContext or skip_prompt_log_probs to True in SamplingParams." + def _dynamic_step_log_probs_bookkeeping(self) -> Tuple[bool, bool]: + """Perform bookkeeping necessary to compute log probs for dynamic batching. - return self.return_log_probs_cuda[:active_request_count].any() - - def _dynamic_step_top_n_logprobs_bookkeeping(self) -> bool: - """Perform bookkeeping necessary to compute top-n log probs for dynamic batching.""" + Returns: + return_log_probs (bool): Whether to return the sampled log_probs. + """ context = self.inference_wrapped_model.inference_context - materialize_only_last_token_logits = context.materialize_only_last_token_logits + active_request_slice = slice(context.paused_request_count, context.total_request_count) - active_request_count = context.total_request_count - context.paused_request_count + return_log_probs = self._request_metadata["return_log_probs"][active_request_slice] + skip_prompt = self._request_metadata["skip_prompt_log_probs"][active_request_slice] + top_n_log_probs = self._request_metadata["top_n_logprobs"][active_request_slice] > 0 - # Check if any request wants prompt top-n logprobs (top_n > 0 AND skip_prompt_log_probs = False) - # Create a copy to avoid modifying the original tensor with in-place operations - to_check = (self.top_n_logprobs_cuda[:active_request_count] > 0).clone() - to_check &= ~self.skip_prompt_log_probs_cuda[:active_request_count] + to_check_prompt = (return_log_probs | top_n_log_probs) & ~skip_prompt - assert not ( - to_check.any() and materialize_only_last_token_logits - ), "Prompt top-n logprobs cannot be calculated if only last token logits are materialized. Set materialize_only_last_token_logits to False in DynamicInferenceContext or set skip_prompt_log_probs to True in SamplingParams." + assert not (to_check_prompt.any() and context.materialize_only_last_token_logits), ( + "Prompt log probs cannot be calculated if only last token logits are materialized. " + "Set materialize_only_last_token_logits to False in DynamicInferenceContext " + "or skip_prompt_log_probs to True in SamplingParams." + ) - # Check if any request has top_n_logprobs > 0 - return (self.top_n_logprobs_cuda[:active_request_count] > 0).any() + return return_log_probs.any(), top_n_log_probs.any() def _dynamic_step_calculate_log_probs(self, logits: Tensor) -> Optional[Tensor]: """Calculate log probs from logits.""" context = self.inference_wrapped_model.inference_context - materialize_only_last_token_logits = context.materialize_only_last_token_logits - active_request_count = context.total_request_count - context.paused_request_count return context.calculate_log_probs( logits, - self.sampled_tokens_cuda[:active_request_count], - only_last_token_logits=materialize_only_last_token_logits, + self._sampled_tokens_cuda[:active_request_count], + only_last_token_logits=context.materialize_only_last_token_logits, ) def _dynamic_step_calculate_top_n_logprobs( @@ -763,19 +693,20 @@ def _dynamic_step_calculate_top_n_logprobs( ) context = self.inference_wrapped_model.inference_context - materialize_only_last_token_logits = context.materialize_only_last_token_logits - active_request_count = context.total_request_count - context.paused_request_count + active_request_slice = slice(context.paused_request_count, context.total_request_count) # Handle decode-only mode (only last token) - if materialize_only_last_token_logits or context.is_decode_only(): + if context.materialize_only_last_token_logits or context.is_decode_only(): # In decode mode or when only last token logits are materialized, # logits already represent only the last tokens log_probs = log_probs_tensor[:active_request_count] top_n_results = {} for req_idx in range(active_request_count): - top_n = int(self.top_n_logprobs_cuda[req_idx].item()) + top_n = int( + self._request_metadata["top_n_logprobs"][active_request_slice][req_idx].item() + ) if top_n > 0: # Get top-n logprobs and indices for this request (single token) top_n_logits = torch.topk(log_probs[req_idx], k=top_n) @@ -789,9 +720,7 @@ def _dynamic_step_calculate_top_n_logprobs( # Note: logits may be padded, so we only take the first active_token_count tokens log_probs = log_probs_tensor[: context.active_token_count] - active_query_lengths = context.request_query_lengths[ - context.paused_request_count : context.total_request_count - ] + active_query_lengths = context.request_query_lengths[active_request_slice] # Split log_probs across request boundaries # log_probs has shape [active_token_count, vocab_size] @@ -799,12 +728,14 @@ def _dynamic_step_calculate_top_n_logprobs( top_n_results = {} for req_idx in range(active_request_count): - top_n = int(self.top_n_logprobs_cuda[req_idx].item()) + top_n = int( + self._request_metadata["top_n_logprobs"][active_request_slice][req_idx].item() + ) if top_n > 0: request_log_probs = log_probs_per_request[ req_idx ] # [num_tokens_for_request, vocab_size] - skip_prompt = bool(self.skip_prompt_log_probs_cuda[req_idx].item()) + skip_prompt = bool(self._request_metadata["skip_prompt_log_probs"][req_idx].item()) # If skip_prompt_log_probs is True, only compute for last token if skip_prompt and request_log_probs.size(0) > 1: @@ -825,9 +756,15 @@ def _dynamic_step_calculate_top_n_logprobs( return top_n_results if top_n_results else None - def _dynamic_step_context_bookkeeping(self, new_sample) -> Dict[str, Tensor]: + def _dynamic_step_context_bookkeeping(self) -> Dict[str, Tensor]: """Update the dynamic inference context after sampling. + Args: + new_sample (Tensor): The newly sampled tokens. + request_metadata (Optional[Dict[str, Tensor]]): An override for the tensors + that manage request metadata, such as sampling parameters. By default, this + metadata is retrieved from the context. + Return: Dict [str, Tensor]: A dictionary containing: active_request_ids (Tensor): Current active request IDs. @@ -835,13 +772,11 @@ def _dynamic_step_context_bookkeeping(self, new_sample) -> Dict[str, Tensor]: finished_request_ids (Tensor): Finished request IDs. """ context = self.inference_wrapped_model.inference_context - active_request_count = context.total_request_count - context.paused_request_count + active_request_slice = slice(context.paused_request_count, context.total_request_count) # Active sequence lengths. - active_request_ids = context.request_ids[ - context.paused_request_count : context.total_request_count - ].long() + active_request_ids = context.request_ids[active_request_slice].long() active_sequence_lengths = context.get_active_sequence_lengths() active_sequence_lengths += 1 # Account for the token we just generated max_sequence_lengths = context.get_max_sequence_lengths() @@ -849,8 +784,8 @@ def _dynamic_step_context_bookkeeping(self, new_sample) -> Dict[str, Tensor]: # Request finished if termination_id or length >= max_sequence_length. # Note: termination_id tensor has per-request termination IDs from mixed sampling active_request_mask = ( - self.sampled_tokens_cuda[:active_request_count] - != self.termination_id_cuda[:active_request_count] + self._sampled_tokens_cuda[:active_request_count] + != self._request_metadata["termination_id"][active_request_slice] ).byte() & torch.less(active_sequence_lengths, max_sequence_lengths).byte() finished_idxs = ( torch.nonzero(active_request_mask == 0, as_tuple=True)[0] + context.paused_request_count @@ -858,7 +793,7 @@ def _dynamic_step_context_bookkeeping(self, new_sample) -> Dict[str, Tensor]: finished_request_ids = context.request_ids[finished_idxs] # New sample gets updated in update_requests, so we pass in a clone - new_sample_copy = new_sample.clone() + new_sample_copy = self._sampled_tokens_cuda[:active_request_count].clone() # Update requests. newly_paused_request_ids = context.update_requests(active_request_mask, new_sample_copy) @@ -888,6 +823,7 @@ async def async_generate_output_tokens_dynamic_batch( cuda_graph_request_count (Optional[int]): Size of cuda graph used for this step. """ context = self.inference_wrapped_model.inference_context + active_request_count = context.total_request_count - context.paused_request_count # No tokens? if context.active_token_count == 0: @@ -910,11 +846,9 @@ async def async_generate_output_tokens_dynamic_batch( # NOTE [TDE]: This will be moved once CPU and GPU methods are separated. await asyncio.sleep(0) + return_log_probs, return_top_n_logprobs = self._dynamic_step_log_probs_bookkeeping() self._dynamic_step_sample_bookkeeping() - new_sample = self._dynamic_step_sample_logits(logits) - - return_log_probs = self._dynamic_step_log_probs_bookkeeping() - return_top_n_logprobs = self._dynamic_step_top_n_logprobs_bookkeeping() + self._dynamic_step_sample_logits(logits) log_probs = None top_n_logprobs = None @@ -928,10 +862,10 @@ async def async_generate_output_tokens_dynamic_batch( if skip_bookkeeping: request_bookkeeping = {} else: - request_bookkeeping = self._dynamic_step_context_bookkeeping(new_sample) + request_bookkeeping = self._dynamic_step_context_bookkeeping() ret = { - "sample": new_sample, + "sample": self._sampled_tokens_cuda[:active_request_count], "log_probs": log_probs, "top_n_logprobs": top_n_logprobs, "cuda_graph_request_count": cuda_graph_request_count, diff --git a/tests/unit_tests/inference/contexts/test_dynamic_context.py b/tests/unit_tests/inference/contexts/test_dynamic_context.py index 456154147f8..2da334191a0 100644 --- a/tests/unit_tests/inference/contexts/test_dynamic_context.py +++ b/tests/unit_tests/inference/contexts/test_dynamic_context.py @@ -509,9 +509,8 @@ def test_add_dummy_requests_parallel_populates_state(self): torch.tensor([2, 1], device='cuda', dtype=torch.int32), ) - termination_idx = DynamicInferenceRequest.get_metadata_labels()["termination_id"] assert torch.equal( - dynamic_context.request_metadata[:2, termination_idx], + dynamic_context.request_metadata["termination_id"][:2], torch.tensor([7.0, 8.0], device='cuda'), ) diff --git a/tests/unit_tests/inference/text_generation_controllers/test_simple_text_generation_controller.py b/tests/unit_tests/inference/text_generation_controllers/test_simple_text_generation_controller.py index 8835b07be07..ebf558d3fa9 100644 --- a/tests/unit_tests/inference/text_generation_controllers/test_simple_text_generation_controller.py +++ b/tests/unit_tests/inference/text_generation_controllers/test_simple_text_generation_controller.py @@ -232,13 +232,16 @@ def detokenize(self, inp, skip_special_tokens=False): ), f"The sampled logits should all be greater than {expected_min_value} but its {sampled_logits}" @pytest.mark.parametrize("backend", ["torch"]) - def test_sample_from_dynamic_logits(self, backend): + @pytest.mark.parametrize("materialize_only_last_token_logits", [True, False]) + def test_sample_from_dynamic_logits( + self, backend: str, materialize_only_last_token_logits: bool + ): batch_size = 12 self.setup_model(torch.float32, batch_size=batch_size, static=False) self.mock_tokenizer.eod = self.vocab_size context = self.text_generation_controller.inference_wrapped_model.inference_context - context.materialize_only_last_token_logits = True + context.materialize_only_last_token_logits = materialize_only_last_token_logits # Prepare sampling params in human-readable format, to aid with test maintenance. sampling_test_cases: List[Tuple[SamplingParams, List[int]]] = [ @@ -258,29 +261,37 @@ def test_sample_from_dynamic_logits(self, backend): rev_sampling_dict[idx] = sampling_params # Prepare metadata for sample bookkeeping. - request_metadata_labels = DynamicInferenceRequest.get_metadata_labels() - request_metadata = torch.empty( - (batch_size, len(request_metadata_labels)), dtype=torch.float32 - ).cuda() - top_k_values = torch.Tensor([s.top_k for s in rev_sampling_dict]).cuda() - request_metadata[:, request_metadata_labels["top_k"]] = top_k_values - top_p_values = torch.Tensor([s.top_p for s in rev_sampling_dict]).cuda() - request_metadata[:, request_metadata_labels["top_p"]] = top_p_values - temp_values = torch.Tensor([s.temperature for s in rev_sampling_dict]).cuda() - request_metadata[:, request_metadata_labels["temperature"]] = temp_values + temp_values = torch.Tensor([s.temperature for s in rev_sampling_dict]) + top_k_values = torch.Tensor([s.top_k for s in rev_sampling_dict]).to(torch.int32) + top_p_values = torch.Tensor([s.top_p for s in rev_sampling_dict]) + request_metadata = { + "temperature": temp_values, + "top_k": top_k_values, + "top_p": top_p_values, + } + self.text_generation_controller._request_metadata = request_metadata + self.text_generation_controller._sampling_backend = backend + + context.padded_active_token_count = batch_size + context.request_query_lengths = torch.ones(batch_size, dtype=torch.int32) + context.paused_request_count = 0 + context.total_request_count = batch_size # Bookkeeping. - self.text_generation_controller._dynamic_step_sample_bookkeeping( - request_metadata=request_metadata - ) + self.text_generation_controller._dynamic_step_sample_bookkeeping() # Sampling. logits = torch.arange(0, self.vocab_size).repeat(batch_size, 1).unsqueeze(0).float().cuda() - sampled_logits = self.text_generation_controller._dynamic_step_sample_logits( - logits, backend=backend - ) + self.text_generation_controller._dynamic_step_sample_logits(logits) + sampled_logits = self.text_generation_controller._sampled_tokens_cuda[:batch_size] vocab_indices = torch.arange(self.vocab_size).cuda() + # Move tensors to GPU for assertion checks. + temp_values = temp_values.cuda() + top_k_values = top_k_values.cuda() + top_p_values = top_p_values.cuda() + + # Assert correct sampled values. top_k_values[top_k_values == 0] = self.vocab_size assert torch.all( sampled_logits >= self.vocab_size - top_k_values @@ -753,21 +764,15 @@ def test_dynamic_top_n_logprobs_calculation( # Prepare sampling params top_n = 5 - request_metadata_labels = DynamicInferenceRequest.get_metadata_labels() - request_metadata = torch.empty( - (batch_size, len(request_metadata_labels)), dtype=torch.float32 - ).cuda() - - # Set top_n_logprobs for all requests - request_metadata[:, request_metadata_labels["top_n_logprobs"]] = top_n - request_metadata[:, request_metadata_labels["skip_prompt_log_probs"]] = float( - skip_prompt_log_probs - ) - - # Bookkeeping - self.text_generation_controller._dynamic_step_sample_bookkeeping( - request_metadata=request_metadata - ) + request_metadata = { + "top_n_logprobs": torch.full((batch_size,), top_n, dtype=torch.int32).cuda(), + "skip_prompt_log_probs": torch.full( + (batch_size,), float(skip_prompt_log_probs), dtype=torch.float32 + ).cuda(), + } + self.text_generation_controller._request_metadata = request_metadata + self.text_generation_controller._active_request_count = batch_size + self.text_generation_controller._active_request_slice = slice(0, batch_size) if materialize_only_last_token_logits: # Decode mode: logits for last tokens only