diff --git a/megatron/core/inference/batch_dimensions_utils.py b/megatron/core/inference/batch_dimensions_utils.py index 4e23151c533..c9474eac5a6 100644 --- a/megatron/core/inference/batch_dimensions_utils.py +++ b/megatron/core/inference/batch_dimensions_utils.py @@ -14,7 +14,7 @@ import torch -from megatron.core.utils import get_pg_size +from megatron.core.utils import get_pg_size, round_up_to_nearest_multiple @dataclass(order=True, frozen=True) @@ -85,6 +85,10 @@ def is_valid( Returns: True if the config is valid, False otherwise """ + # A dimension with no tokens serves no requests. + if self.token_count <= 0: + return False + # Check if total requests exceed maximum if self.prefill_req_count + self.decode_req_count > max_requests: return False @@ -269,7 +273,9 @@ def _calculate_cuda_graph_token_counts( ) # Align each entry to TP size cuda_graph_token_counts = list( - dict.fromkeys(math.ceil(s / tp_size) * tp_size for s in cuda_graph_token_counts) + dict.fromkeys( + round_up_to_nearest_multiple(s, tp_size) for s in cuda_graph_token_counts + ) ) # Clamp to max tokens cuda_graph_token_counts = [ @@ -291,7 +297,7 @@ def _calculate_cuda_graph_token_counts( math.ceil(int(cuda_graph_step_size) / CUDAGraphBatchDimensionBuilder.CUDA_GRAPH_ROUNDER) ) # Make sure divisible by TP size - cuda_graph_step_size = math.ceil(cuda_graph_step_size / tp_size) * tp_size + cuda_graph_step_size = round_up_to_nearest_multiple(cuda_graph_step_size, tp_size) # round down cuda graph max tokens to be multiple of TP size cuda_graph_max_tokens = (cuda_graph_max_tokens // tp_size) * tp_size diff --git a/megatron/core/inference/engines/dynamic_engine.py b/megatron/core/inference/engines/dynamic_engine.py index 97f3e098c2b..f96235db0c0 100644 --- a/megatron/core/inference/engines/dynamic_engine.py +++ b/megatron/core/inference/engines/dynamic_engine.py @@ -50,7 +50,7 @@ unset_inference_cuda_graphed_iteration_for_ep_inference, ) from megatron.core.process_groups_config import ProcessGroupCollection -from megatron.core.transformer.cuda_graphs import delete_cuda_graphs +from megatron.core.transformer.cuda_graphs import delete_cuda_graphs, graph_capture from megatron.core.transformer.enums import CudaGraphScope from megatron.core.transformer.moe.router_replay import RouterReplay, RouterReplayAction from megatron.core.utils import ( @@ -63,7 +63,9 @@ internal_api, nvtx_range_pop, nvtx_range_push, + round_up_to_nearest_multiple, trace_async_exceptions, + unwrap_model, ) from .async_zmq_communicator import AsyncZMQCommunicator @@ -365,6 +367,21 @@ def create_cuda_graphs(self, reset_context: bool = True): unwrapped_model = controller.inference_wrapped_model.model set_inference_cuda_graphed_iteration_for_ep_inference(unwrapped_model) + # MTP warmup preparation: capture MTP CUDA graphs alongside the + # decoder graphs within the same loop rather than in a separate pass. + unwrapped = unwrap_model(controller.inference_wrapped_model.model) + mtp_warmup_enabled = ( + controller.num_mtp_heads > 0 + and (controller.num_speculative_tokens or 0) > 0 + and hasattr(unwrapped, 'mtp') + ) + if mtp_warmup_enabled: + tp_size = get_pg_size(controller.inference_wrapped_model.tp_group) + sp_enabled = model_config.sequence_parallel and tp_size > 1 + mtp_pass_depth = not unwrapped.mtp.mtp_use_repeated_layer + mtp_warmup_depths = range(controller._num_mtp_depths) if mtp_pass_depth else [None] + mtp_seen_batch_sizes = set() + tbar = enumerate(context.cuda_graph_batch_dimensions_list) if HAVE_TQDM: tbar = tqdm(tbar, total=len(context.cuda_graph_batch_dimensions_list)) @@ -390,12 +407,39 @@ def create_cuda_graphs(self, reset_context: bool = True): # Forward pass -> logits. controller._dynamic_step_forward_logits(input_ids, position_ids) + # MTP CUDA graph warmup for this batch dimension. + if mtp_warmup_enabled: + n = cuda_graph_batch_dimension.req_count + if sp_enabled: + n = round_up_to_nearest_multiple(n, tp_size) + if n > 0 and n not in mtp_seen_batch_sizes: + mtp_seen_batch_sizes.add(n) + device = torch.cuda.current_device() + batch_dim = n // tp_size if sp_enabled else n + # Use zeros (not empty) — garbage token IDs cause OOB embedding lookups during graph capture/replay. + for depth in mtp_warmup_depths: + with graph_capture(): + unwrapped.compute_mtp_single_step( + hidden_states=torch.zeros( + (batch_dim, 1, model_config.hidden_size), + device=device, + dtype=model_config.params_dtype, + ), + next_token_ids=torch.zeros((1, n), device=device, dtype=torch.long), + position_ids=torch.zeros((1, n), device=device, dtype=torch.int64), + depth=depth, + ) + context.reset() # Disable inference dispatcher after graph capture if is_inference_optimized_ep: unset_inference_cuda_graphed_iteration_for_ep_inference(unwrapped_model) + if mtp_warmup_enabled and mtp_seen_batch_sizes: + controller.has_mtp_cuda_graphs = True + logging.info("> MTP CUDA graph warmup: %d batch size(s)", len(mtp_seen_batch_sizes)) + # Memory usage. time_end = time.time() mem_stats_end = torch.cuda.memory_stats() diff --git a/megatron/core/inference/text_generation_controllers/mtp_utils_pytorch.py b/megatron/core/inference/text_generation_controllers/mtp_utils_pytorch.py new file mode 100644 index 00000000000..59bad67d70a --- /dev/null +++ b/megatron/core/inference/text_generation_controllers/mtp_utils_pytorch.py @@ -0,0 +1,244 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import torch + + +def rewind_kv_cache( + accepted_counts, + prefill_status, + last_kv_block_offset, + kv_length_offsets, + kv_block_counts, + last_kv_block_id, + kv_block_ids, + num_speculative_tokens, + block_size_tokens, + num_active_requests=None, +): + """Update the KV cache bookkeeping for speculative decoding. + + After forward pass with speculative tokens, some tokens may be rejected. + This function "rewinds" the KV cache bookkeeping to reflect only the accepted tokens. + + When speculative tokens are rejected, we need to: + 1. Update kv_length_offsets (total sequence length) + 2. Update last_kv_block_offset (position within last block) + 3. If rewinding crosses a block boundary: + - Reduce kv_block_counts + - Update last_kv_block_id to point to the previous block + - Clear the entry in kv_block_ids for the released block + + Mutates the input tensors in-place. + + Returns (blocks_to_release, remove_mask). + """ + N = accepted_counts.shape[0] + if num_active_requests is None: + num_active_requests = N + + blocks_to_release = torch.empty_like(last_kv_block_id) + remove_mask = torch.empty(N, device=accepted_counts.device, dtype=torch.bool) + + for i in range(N): + if i >= num_active_requests: + blocks_to_release[i] = 0 + remove_mask[i] = False + continue + + accepted = accepted_counts[i].item() + prefill = prefill_status[i].item() + last_offset = last_kv_block_offset[i].item() + kv_length = kv_length_offsets[i].item() + block_count = kv_block_counts[i].item() + last_block = last_kv_block_id[i].item() + + # Number of tokens to rewind (rejected speculative tokens). + # For prefill requests, no speculative tokens were forwarded through the model, + # so there is nothing to rewind. + num_to_rewind = 0 if prefill == 1 else num_speculative_tokens - accepted + + # Save the original offset BEFORE modifying to correctly detect block boundary crossing. + # A request crosses back to a previous block if: original_offset - num_to_rewind < 0 + diff = last_offset - num_to_rewind + remove = diff < 0 + + # Update the offsets + new_offset = diff % block_size_tokens + last_kv_block_offset[i] = new_offset + kv_length_offsets[i] = kv_length - num_to_rewind + + # For requests that crossed back to a previous block, we need to: + # 1. Reduce the block count by 1 + # 2. Get the block ID to release (current last_kv_block_id) + # 3. Update last_kv_block_id to point to the previous block + # 4. Clear the entry in kv_block_ids for the released block + # 5. Release the block back to the allocator + blocks_to_release[i] = last_block + + # Reduce block counts for requests that crossed back + new_block_count = block_count - 1 if remove else block_count + kv_block_counts[i] = new_block_count + + # Update last_kv_block_id to point to the previous block (at index new_count - 1) + prev_idx = max(new_block_count - 1, 0) + prev_block_id = kv_block_ids[i, prev_idx].item() + last_kv_block_id[i] = prev_block_id if remove else last_block + + # Clear the released block entry (at index new_count, which was the old last block) + scatter_idx = min(new_block_count, kv_block_ids.shape[1] - 1) + if remove: + kv_block_ids[i, scatter_idx] = -1 + + remove_mask[i] = remove + + return blocks_to_release, remove_mask + + +# pylint: disable=line-too-long +def verify_speculative_tokens( + input_tokens, output_tokens, num_decode_requests, num_prefill_requests, num_speculative_tokens +): + """Verify speculative tokens against input tokens and compute acceptance. + + Creates an accepted tokens mask where: + - For prefill requests, the token is always accepted. + - For decode requests, the first token (base token) is always accepted, then we compare + sampled tokens with input tokens and accept consecutive matches. + Then finds the index of the last accepted token per request. + + Example (assume 1, 2, and 0 spec tokens are accepted in the first 3 decode requests): + input_tokens_required: [ a5 a6s a7s | b3 b4s b5s | c6 c7s c8s | d2 | e4 ] # Size 11 + Output tokens [ a6o a7o a8o | b40 b5o b6o | c7o c8o c9o | d3o | e5o ] + Output tokens right shift [ d3o a6o a7o | a8o b40 b5o | b6o c7o c8o | c9o | d3o ] + Accepted tokens mask [ 1 1 0 | 1 1 1 | 1 0 0 | 1 | 1 ] + Last one indices [ 1 | 5 | 6 | 9 | 10 ] + + Returns: + tuple: (last_one_indices, accepted_tokens_mask, input_tokens) where + last_one_indices contains the index of the last accepted token per request. + """ + if input_tokens.ndim == 2: + input_tokens = input_tokens.squeeze(0) + + stride = num_speculative_tokens + 1 + active_request_count = num_decode_requests + num_prefill_requests + decode_len = num_decode_requests * stride + + # Initialize mask with False to prevent boundary bleed + accepted_tokens_mask = torch.zeros_like(input_tokens, dtype=torch.bool) + + # Safe decode token verification without cross-batch boundary contamination + decode_mask_2d = None + if num_decode_requests > 0: + decode_inputs = input_tokens[:decode_len].reshape(num_decode_requests, stride) + decode_outputs = output_tokens[:decode_len].reshape(num_decode_requests, stride) + + # Shift outputs right by 1 *within* each request to align sampled tokens with input targets + decode_outputs_shifted = decode_outputs.roll(1, dims=1) + decode_mask_2d = decode_inputs == decode_outputs_shifted + # The first token (base token) is always accepted + decode_mask_2d[:, 0] = True + # Enforce consecutive acceptance: cummin propagates False to the right + decode_mask_2d = decode_mask_2d.cummin(dim=1).values + accepted_tokens_mask[:decode_len] = decode_mask_2d.flatten() + + # Make all prefill tokens accepted + if num_prefill_requests > 0: + accepted_tokens_mask[decode_len:] = True + + last_one_indices = torch.full( + (active_request_count,), -1, device=input_tokens.device, dtype=torch.long + ) + + if num_decode_requests > 0: + # Summing the consecutive mask gives the count; subtract 1 for the local index + local_last_indices = decode_mask_2d.sum(dim=1) - 1 + row_offsets = torch.arange(num_decode_requests, device=input_tokens.device) * stride + last_one_indices[:num_decode_requests] = row_offsets + local_last_indices + + if num_prefill_requests > 0: + prefill_valid = torch.nonzero(accepted_tokens_mask[decode_len:]).squeeze(-1) + decode_len + last_one_indices[num_decode_requests:] = prefill_valid + + return last_one_indices, accepted_tokens_mask, input_tokens + + +# pylint: disable=line-too-long +def prepare_next_forward_pass( + num_decode_requests, + output_tokens, + required_logit_indices, + last_one_indices, + accepted_tokens_mask, + input_tokens, + sampled_tokens_buf, + last_accepted_seq_buf, + accepted_tokens_per_request, + accepted_token_counts, + num_speculative_tokens, +): + """Prepare data for the next forward pass after speculative token verification. + + For each active request: + - Store the final sampled tokens for the next forward pass. + - Store the last accepted positions in the packed sequence for serial + MTP computation after verification. + + For decode requests, extract accepted tokens and counts: + input_tokens_required: [ a5 a6s a7s | b3 b4s b5s | c6 c7s c8s | d2 | e4 ] + Accepted tokens mask [ 1 1 0 | 1 1 1 | 1 0 0 | 1 | 1 ] + Accepted tokens [ [a6s -1] | [b4s b5s] | [-1 -1] ] # Only decode requests (prefill defaults to -1) + Accepted token counts [ 1 | 2 | 0 ] # Prefill defaults to 0 + + Writes results into the pre-allocated buffers provided by the caller. + """ + active_request_count = last_one_indices.shape[0] + stride = num_speculative_tokens + 1 + + for pid in range(active_request_count): + idx = last_one_indices[pid].item() + + # Store the final sampled tokens for the next forward pass. + sampled_tokens_buf[pid] = output_tokens[idx] + + # Store the last accepted positions in the packed sequence for serial + # MTP computation after verification. + last_accepted_seq_buf[pid] = required_logit_indices[idx] + + # Extract accepted tokens and counts for decode requests. + # For prefill it is always set to 1. For decode, the first token is always accepted, + # then we compare with input tokens and accept the next tokens if its a match. + if pid < num_decode_requests: + base = pid * stride + # Skip the first token of every decode request (i.e a5, b3, c6) + for s in range(num_speculative_tokens): + pos = base + 1 + s + if accepted_tokens_mask[pos]: + accepted_tokens_per_request[pid, s] = input_tokens[pos] + else: + accepted_tokens_per_request[pid, s] = -1 + + count = 0 + for s in range(num_speculative_tokens): + if accepted_tokens_per_request[pid, s].item() != -1: + count += 1 + accepted_token_counts[pid] = count + + +def mamba_state_selective_copy( + intermediate_states, current_states, prefill_status, state_idx, accepted_counts, num_layers +): + """Mamba speculative rewind state update. + + For each decode request, copies + `intermediate[layer, slot, accepted_count, ...]` → + `current[layer, slot, ...]` for every Mamba layer. + """ + N = prefill_status.shape[0] + for i in range(N): + if prefill_status[i].item() == 1: + continue + slot = state_idx[i].item() + accepted = accepted_counts[i].item() + for layer in range(num_layers): + current_states[layer, slot] = intermediate_states[layer, slot, accepted] diff --git a/megatron/core/inference/text_generation_controllers/mtp_utils_triton.py b/megatron/core/inference/text_generation_controllers/mtp_utils_triton.py new file mode 100644 index 00000000000..37ff55c1e99 --- /dev/null +++ b/megatron/core/inference/text_generation_controllers/mtp_utils_triton.py @@ -0,0 +1,456 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import math + +import torch + +try: + import triton + import triton.language as tl + + HAVE_TRITON = True +except ImportError: + from unittest.mock import MagicMock + + from megatron.core.utils import null_decorator + + triton = MagicMock() + triton.jit = null_decorator + tl = MagicMock() + HAVE_TRITON = False + + +# --------------------------------------------------------------------------- +# Kernel 1: KV-cache rewind for speculative decoding +# --------------------------------------------------------------------------- +@triton.jit +def _rewind_kv_cache_kernel( + # Per-request input (read-only) + ACCEPTED_COUNTS_PTR, + PREFILL_STATUS_PTR, + # Per-request state (read-write, updated in-place) + LAST_KV_BLOCK_OFFSET_PTR, + KV_LENGTH_OFFSETS_PTR, + KV_BLOCK_COUNTS_PTR, + LAST_KV_BLOCK_ID_PTR, + # 2-D table [N, max_blocks] (read-write) + KV_BLOCK_IDS_PTR, + # Per-request outputs + BLOCKS_TO_RELEASE_PTR, + REMOVE_MASK_PTR, + # Strides / limits + kv_block_ids_stride, + max_blocks_minus_1, + num_active_requests, + # Compile-time constants + NUM_SPEC_TOKENS: tl.constexpr, + BLOCK_SIZE_TOKENS: tl.constexpr, +): + """Rewind KV-cache bookkeeping for one request after speculative verification. + + Grid: may be padded beyond active requests for CUDA-graph compatibility. + Each program handles exactly one request. Programs with + `pid >= num_active_requests` are padding and produce safe no-op outputs. + """ + pid = tl.program_id(0) + + # Padding programs: write safe defaults and skip all state mutation. + if pid >= num_active_requests: + tl.store(BLOCKS_TO_RELEASE_PTR + pid, 0) + tl.store(REMOVE_MASK_PTR + pid, False) + return + + # --- Load per-request scalars --- + accepted = tl.load(ACCEPTED_COUNTS_PTR + pid) + prefill = tl.load(PREFILL_STATUS_PTR + pid) + last_offset = tl.load(LAST_KV_BLOCK_OFFSET_PTR + pid) + kv_length = tl.load(KV_LENGTH_OFFSETS_PTR + pid) + block_count = tl.load(KV_BLOCK_COUNTS_PTR + pid) + last_block_id = tl.load(LAST_KV_BLOCK_ID_PTR + pid) + + # --- Compute rewind (zero for prefill requests) --- + num_to_rewind = tl.where(prefill == 1, 0, NUM_SPEC_TOKENS - accepted) + diff = last_offset - num_to_rewind + remove = diff < 0 + + # Python-style modulo: ((diff % M) + M) % M to handle negative diff + new_offset = ((diff % BLOCK_SIZE_TOKENS) + BLOCK_SIZE_TOKENS) % BLOCK_SIZE_TOKENS + tl.store(LAST_KV_BLOCK_OFFSET_PTR + pid, new_offset) + tl.store(KV_LENGTH_OFFSETS_PTR + pid, kv_length - num_to_rewind) + + # Save current last block id (will be released by caller if remove is True) + tl.store(BLOCKS_TO_RELEASE_PTR + pid, last_block_id) + + # Decrement block count when a block boundary was crossed + new_block_count = tl.where(remove, block_count - 1, block_count) + tl.store(KV_BLOCK_COUNTS_PTR + pid, new_block_count) + + # Gather previous block id from the 2-D table + kv_row_base = pid.to(tl.int64) * kv_block_ids_stride + prev_idx = tl.maximum(new_block_count - 1, 0) + prev_block_id = tl.load(KV_BLOCK_IDS_PTR + kv_row_base + prev_idx) + + # Conditionally update last block id + tl.store(LAST_KV_BLOCK_ID_PTR + pid, tl.where(remove, prev_block_id, last_block_id)) + + # Clear released block entry via scatter + scatter_idx = tl.minimum(new_block_count, max_blocks_minus_1) + current_val = tl.load(KV_BLOCK_IDS_PTR + kv_row_base + scatter_idx) + tl.store(KV_BLOCK_IDS_PTR + kv_row_base + scatter_idx, tl.where(remove, -1, current_val)) + + # Output remove mask for the caller (to release blocks outside this kernel) + tl.store(REMOVE_MASK_PTR + pid, remove) + + +def rewind_kv_cache( + accepted_counts, + prefill_status, + last_kv_block_offset, + kv_length_offsets, + kv_block_counts, + last_kv_block_id, + kv_block_ids, + num_speculative_tokens, + block_size_tokens, + num_active_requests=None, +): + """Launch the KV-cache rewind Triton kernel. + + Args: + num_active_requests: Number of real (non-padding) requests. When the + grid is padded beyond this count, the kernel skips padding + programs so stale data in padding slots cannot corrupt + bookkeeping. Defaults to `accepted_counts.shape[0]` (no + padding). + + Returns: + (blocks_to_release, remove_mask) — same semantics as the original + torch.compile'd `_rewind_kv_cache` (KV-cache portion only; Mamba + state updates are handled separately by the caller). + """ + N = accepted_counts.shape[0] + if num_active_requests is None: + num_active_requests = N + if N == 0: + return ( + torch.empty(0, device=accepted_counts.device, dtype=last_kv_block_id.dtype), + torch.empty(0, device=accepted_counts.device, dtype=torch.bool), + ) + + blocks_to_release = torch.empty_like(last_kv_block_id) + remove_mask = torch.empty(N, device=accepted_counts.device, dtype=torch.bool) + + _rewind_kv_cache_kernel[(N,)]( + accepted_counts, + prefill_status, + last_kv_block_offset, + kv_length_offsets, + kv_block_counts, + last_kv_block_id, + kv_block_ids, + blocks_to_release, + remove_mask, + kv_block_ids_stride=kv_block_ids.stride(0), + max_blocks_minus_1=kv_block_ids.shape[1] - 1, + num_active_requests=num_active_requests, + NUM_SPEC_TOKENS=num_speculative_tokens, + BLOCK_SIZE_TOKENS=block_size_tokens, + ) + return blocks_to_release, remove_mask + + +# --------------------------------------------------------------------------- +# Kernel 2: Verify speculative tokens +# --------------------------------------------------------------------------- +@triton.jit +def _verify_speculative_tokens_kernel( + INPUT_TOKENS_PTR, + OUTPUT_TOKENS_PTR, + # Outputs + ACCEPTED_MASK_PTR, + LAST_ONE_INDICES_PTR, + # Runtime scalars + num_decode_requests, + decode_len, + # Compile-time constants + STRIDE: tl.constexpr, # num_speculative_tokens + 1 + BLOCK_SIZE: tl.constexpr, # next_power_of_2(STRIDE) +): + """Verify speculative tokens for one request. + + Grid: (active_request_count,) + Programs 0..num_decode_requests-1 handle decode requests. + Programs num_decode_requests..end handle prefill requests. + """ + pid = tl.program_id(0) + + if pid < num_decode_requests: + base = pid * STRIDE + offsets = tl.arange(0, BLOCK_SIZE) + valid = offsets < STRIDE + + input_toks = tl.load(INPUT_TOKENS_PTR + base + offsets, mask=valid, other=0) + + # Build shifted output: shifted[i] = output[i-1]. + # Position 0 uses a dummy load (always accepted regardless). + safe_shifted = tl.where(offsets > 0, offsets - 1, 0) + shifted_output = tl.load(OUTPUT_TOKENS_PTR + base + safe_shifted, mask=valid, other=0) + + # First token is always accepted; rest must match shifted output. + match = tl.where(offsets == 0, 1, (input_toks == shifted_output).to(tl.int32)) + match = tl.where(valid, match, 0) + + # Consecutive acceptance via cumulative-sum trick: + # accepted[i] iff cumsum(match)[i] == i + 1 + cumsum = tl.cumsum(match, axis=0) + accepted = (cumsum == (offsets + 1)) & valid + + tl.store(ACCEPTED_MASK_PTR + base + offsets, accepted, mask=valid) + + accepted_count = tl.sum(accepted.to(tl.int32)) + tl.store(LAST_ONE_INDICES_PTR + pid, (base + accepted_count - 1).to(tl.int64)) + else: + # Prefill request — single token, always accepted + prefill_idx = decode_len + (pid - num_decode_requests) + tl.store(ACCEPTED_MASK_PTR + prefill_idx, 1) + tl.store(LAST_ONE_INDICES_PTR + pid, prefill_idx.to(tl.int64)) + + +def verify_speculative_tokens( + input_tokens, output_tokens, num_decode_requests, num_prefill_requests, num_speculative_tokens +): + """Launch the speculative-token verification Triton kernel. + + Returns: + (last_one_indices, accepted_tokens_mask, input_tokens) + matching the original `_verify_speculative_tokens` signature. + """ + if input_tokens.ndim == 2: + input_tokens = input_tokens.squeeze(0) + + device = input_tokens.device + active_request_count = num_decode_requests + num_prefill_requests + stride = num_speculative_tokens + 1 + decode_len = num_decode_requests * stride + + accepted_tokens_mask = torch.zeros_like(input_tokens, dtype=torch.bool) + last_one_indices = torch.full((active_request_count,), -1, device=device, dtype=torch.long) + + if active_request_count > 0: + block_size = triton.next_power_of_2(stride) + _verify_speculative_tokens_kernel[(active_request_count,)]( + input_tokens, + output_tokens, + accepted_tokens_mask, + last_one_indices, + num_decode_requests=num_decode_requests, + decode_len=decode_len, + STRIDE=stride, + BLOCK_SIZE=block_size, + ) + + return last_one_indices, accepted_tokens_mask, input_tokens + + +# --------------------------------------------------------------------------- +# Kernel 3: Prepare speculative tokens for next forward pass +# --------------------------------------------------------------------------- +@triton.jit +def _prepare_next_forward_pass_kernel( + OUTPUT_TOKENS_PTR, + REQUIRED_LOGIT_INDICES_PTR, + LAST_ONE_INDICES_PTR, + INPUT_TOKENS_PTR, + ACCEPTED_MASK_PTR, + # Outputs + SAMPLED_TOKENS_OUT_PTR, + LAST_ACCEPTED_SEQ_OUT_PTR, + ACCEPTED_TOKENS_OUT_PTR, + ACCEPTED_COUNTS_OUT_PTR, + # Strides + accepted_tokens_out_stride, + # Runtime scalars + num_decode_requests, + # Compile-time constants + STRIDE: tl.constexpr, # num_speculative_tokens + 1 + NUM_SPEC_TOKENS: tl.constexpr, + SPEC_BLOCK_SIZE: tl.constexpr, # next_power_of_2(NUM_SPEC_TOKENS) +): + """Gather final tokens and extract accepted speculative tokens per request. + + Grid: (active_request_count,) + """ + pid = tl.program_id(0) + + # --- Gather final sampled token and sequence index for every request --- + idx = tl.load(LAST_ONE_INDICES_PTR + pid) + tl.store(SAMPLED_TOKENS_OUT_PTR + pid, tl.load(OUTPUT_TOKENS_PTR + idx)) + tl.store(LAST_ACCEPTED_SEQ_OUT_PTR + pid, tl.load(REQUIRED_LOGIT_INDICES_PTR + idx)) + + # --- For decode requests: extract accepted tokens and count --- + if pid < num_decode_requests: + base = pid * STRIDE + spec_offsets = tl.arange(0, SPEC_BLOCK_SIZE) + spec_valid = spec_offsets < NUM_SPEC_TOKENS + token_positions = base + 1 + spec_offsets # skip first (base) token + + tokens = tl.load(INPUT_TOKENS_PTR + token_positions, mask=spec_valid, other=0) + mask_val = tl.load(ACCEPTED_MASK_PTR + token_positions, mask=spec_valid, other=0) + accepted = mask_val != 0 + + result = tl.where(accepted & spec_valid, tokens, -1) + + out_base = pid.to(tl.int64) * accepted_tokens_out_stride + tl.store(ACCEPTED_TOKENS_OUT_PTR + out_base + spec_offsets, result, mask=spec_valid) + + count = tl.sum((accepted & spec_valid).to(tl.int64)) + tl.store(ACCEPTED_COUNTS_OUT_PTR + pid, count) + + +def prepare_next_forward_pass( + num_decode_requests, + output_tokens, + required_logit_indices, + last_one_indices, + accepted_tokens_mask, + input_tokens, + sampled_tokens_buf, + last_accepted_seq_buf, + accepted_tokens_per_request, + accepted_token_counts, + num_speculative_tokens, +): + """Launch the prepare-next-forward-pass Triton kernel. + + Writes results into the pre-allocated buffers provided by the caller. + """ + active_request_count = last_one_indices.shape[0] + if active_request_count == 0: + return + + stride = num_speculative_tokens + 1 + spec_block_size = triton.next_power_of_2(num_speculative_tokens) + + _prepare_next_forward_pass_kernel[(active_request_count,)]( + output_tokens, + required_logit_indices, + last_one_indices, + input_tokens, + accepted_tokens_mask, + sampled_tokens_buf, + last_accepted_seq_buf, + accepted_tokens_per_request, + accepted_token_counts, + accepted_tokens_out_stride=accepted_tokens_per_request.stride(0), + num_decode_requests=num_decode_requests, + STRIDE=stride, + NUM_SPEC_TOKENS=num_speculative_tokens, + SPEC_BLOCK_SIZE=spec_block_size, + ) + + +# --------------------------------------------------------------------------- +# Kernel 4: Mamba state selective copy (eliminates temporary allocations) +# --------------------------------------------------------------------------- +@triton.jit +def _mamba_state_selective_copy_kernel( + # Source: intermediate states [L, M, S+1, *state_shape] + SRC_PTR, + # Destination: current states [L, M, *state_shape] + DST_PTR, + # Per-request index arrays + PREFILL_STATUS_PTR, # [N] 0=decode, 1=prefill + STATE_IDX_PTR, # [N] maps request → mamba state slot + ACCEPTED_PTR, # [N] accepted token index per request + # Strides (in elements) + src_stride_layer, + src_stride_slot, + src_stride_spec, + dst_stride_layer, + dst_stride_slot, + # Data size + STATE_SIZE, + # Compile-time + BLOCK_SIZE: tl.constexpr, +): + """Copy intermediate Mamba state to current state for decode requests. + + Grid: (N, L, num_chunks) + - dim 0: active request index + - dim 1: mamba layer index + - dim 2: chunk of the flattened state vector + + No-op for prefill requests. + """ + pid_req = tl.program_id(0) + pid_layer = tl.program_id(1) + pid_chunk = tl.program_id(2) + + # Skip prefill requests immediately. + prefill = tl.load(PREFILL_STATUS_PTR + pid_req) + if prefill == 1: + return + + state_idx = tl.load(STATE_IDX_PTR + pid_req).to(tl.int64) + accepted = tl.load(ACCEPTED_PTR + pid_req).to(tl.int64) + + chunk_start = pid_chunk * BLOCK_SIZE + offsets = tl.arange(0, BLOCK_SIZE) + elem_offsets = chunk_start + offsets + mask = elem_offsets < STATE_SIZE + + src_base = ( + pid_layer.to(tl.int64) * src_stride_layer + + state_idx * src_stride_slot + + accepted * src_stride_spec + ) + dst_base = pid_layer.to(tl.int64) * dst_stride_layer + state_idx * dst_stride_slot + + data = tl.load(SRC_PTR + src_base + elem_offsets, mask=mask) + tl.store(DST_PTR + dst_base + elem_offsets, data, mask=mask) + + +def mamba_state_selective_copy( + intermediate_states, current_states, prefill_status, state_idx, accepted_counts, num_layers +): + """Copy accepted intermediate Mamba states to current states in-place. + + For each decode request, copies + `intermediate[layer, slot, accepted_count, ...]` → + `current[layer, slot, ...]` for every Mamba layer. + + Args: + intermediate_states: `(L, M, S+1, *state_shape)` — intermediate buffer. + current_states: `(L, M, *state_shape)` — current state buffer (updated in-place). + prefill_status: `(N,)` int tensor — 0 for decode, 1 for prefill. + state_idx: `(N,)` int tensor — mamba state slot index per request. + accepted_counts: `(N,)` int tensor — accepted token index per request. + num_layers: number of Mamba layers (first dim of the state tensors). + """ + N = prefill_status.shape[0] + if N == 0: + return + + # The state vector to copy per (layer, request) is the product of all + # trailing dimensions after the speculative-token axis. + # intermediate shape: (L, M, S+1, *state_shape) → state_size = prod(state_shape) + state_size = math.prod(intermediate_states.shape[3:]) + + BLOCK_SIZE = 1024 + num_chunks = triton.cdiv(state_size, BLOCK_SIZE) + grid = (N, num_layers, num_chunks) + + _mamba_state_selective_copy_kernel[grid]( + intermediate_states, + current_states, + prefill_status, + state_idx, + accepted_counts, + src_stride_layer=intermediate_states.stride(0), + src_stride_slot=intermediate_states.stride(1), + src_stride_spec=intermediate_states.stride(2), + dst_stride_layer=current_states.stride(0), + dst_stride_slot=current_states.stride(1), + STATE_SIZE=state_size, + BLOCK_SIZE=BLOCK_SIZE, + ) 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 2798533e783..993d05afbe0 100644 --- a/megatron/core/inference/text_generation_controllers/text_generation_controller.py +++ b/megatron/core/inference/text_generation_controllers/text_generation_controller.py @@ -40,6 +40,9 @@ get_asyncio_loop, get_model_config, get_pg_size, + nvtx_range_pop, + nvtx_range_push, + round_up_to_nearest_multiple, unwrap_model, ) @@ -52,6 +55,12 @@ HAVE_TE = False from megatron.core.inference.batch_dimensions_utils import InferenceBatchDimensions +from megatron.core.inference.text_generation_controllers.mtp_utils_triton import ( + mamba_state_selective_copy, + prepare_next_forward_pass, + rewind_kv_cache, + verify_speculative_tokens, +) # pylint: disable=line-too-long @@ -91,6 +100,7 @@ def __init__(self, inference_wrapped_model: AbstractModelInferenceWrapper, token self.sampling_rng = torch.Generator(device=torch.cuda.current_device()) self.num_mtp_heads = self._get_mtp_num_heads() + self.has_mtp_cuda_graphs = False self.sampling_rng.manual_seed(self.model_config.inference_sampling_seed) if ( @@ -137,12 +147,6 @@ def _init_dynamic_sampling_tensors(self): self._sampling_backend = "torch" self._sampled_tokens_cuda = torch.empty(max_requests, dtype=torch.int64, device=device) - # Speculative tokens tensor will be allocated later when num_speculative_tokens is set by the engine - self._accepted_tokens_per_request = None - # MTP tensor will be allocated later when num_speculative_tokens is set by the engine - self._sampled_mtp_tokens_cuda = None - # Last accepted sequence indices for serial MTP computation - self._last_accepted_seq_indices = None # Keep track of request metadata. self._request_metadata: Dict[str, Tensor] = {} @@ -158,23 +162,49 @@ def _init_dynamic_sampling_tensors(self): if self._sampling_backend == "torch": self._torch_sampling_buckets: List[Tuple] = [] - self._init_mtp_sampling_tensor() + # Cache values that are constant across inference steps. + self._unwrapped_model = unwrap_model(self.inference_wrapped_model.model) + self._is_last_pp_stage = is_pipeline_last_stage(self.pp_group) + self._tp_size = get_pg_size(self.inference_wrapped_model.tp_group) + self._sp_enabled = self.model_config.sequence_parallel and self._tp_size > 1 - def _init_mtp_sampling_tensor(self): - """Initialize the MTP sampling tensor after num_speculative_tokens is set.""" - if self.num_speculative_tokens is not None and self.num_speculative_tokens > 0: - context = self.inference_wrapped_model.inference_context - max_requests = context.max_requests - device = torch.cuda.current_device() - self._sampled_mtp_tokens_cuda = torch.empty( - [self.num_speculative_tokens, max_requests], dtype=torch.int64, device=device - ) - self._accepted_tokens_per_request = ( - torch.ones( - [max_requests, self.num_speculative_tokens], dtype=torch.int64, device=device - ) - * -1 + self._init_mtp_sampling_tensors() + + def _init_mtp_sampling_tensors(self): + """Pre-allocate MTP sampling tensors. + + Addresses must be stable across steps for CUDA graph capture. + """ + if not self.num_speculative_tokens: + self._sampled_mtp_tokens_cuda = None + self._accepted_tokens_per_request = None + self._last_accepted_seq_indices = None + return + + context = self.inference_wrapped_model.inference_context + max_requests = context.max_requests + device = torch.cuda.current_device() + self._sampled_mtp_tokens_cuda = torch.empty( + [self.num_speculative_tokens, max_requests], dtype=torch.int64, device=device + ) + self._accepted_tokens_per_request = ( + torch.ones( + [max_requests, self.num_speculative_tokens], dtype=torch.int64, device=device ) + * -1 + ) + self._accepted_token_counts_per_request = torch.zeros( + max_requests, dtype=torch.int64, device=device + ) + self._last_accepted_seq_indices_buf = torch.empty( + max_requests, dtype=torch.int64, device=device + ) + self._last_accepted_seq_indices = None + self._num_mtp_depths = min(self.num_speculative_tokens, self.num_mtp_heads) + self._mtp_token_ids_buf = torch.empty([1, max_requests], dtype=torch.int64, device=device) + self._mtp_position_ids_buf = torch.empty( + [1, max_requests], dtype=torch.int64, device=device + ) @staticmethod def tokenize_prompt(tokenizer, prompt: str, add_BOS: bool = False) -> List[int]: @@ -582,6 +612,25 @@ def _dynamic_step_context_init( is_expert_parallel_dummy_cuda_graph_step=is_dummy_forward, ) + # Derive the MTP padded batch size from the existing padded graph dimensions. + # For MoE models this is post EP sync. In eager mode MTP uses locally SP-aligned + # batch size instead. + if self.has_mtp_cuda_graphs and context.using_cuda_graph_this_step(): + self._mtp_resolved_padded_count = context.padded_batch_dimensions.req_count + if self._sp_enabled: + self._mtp_resolved_padded_count = round_up_to_nearest_multiple( + self._mtp_resolved_padded_count, self._tp_size + ) + else: + self._mtp_resolved_padded_count = None + + # Tell the model whether to use MTP CUDA graphs this step. When the + # main model falls back to eager mode, MTP must also run eagerly across + # all EP ranks — otherwise some ranks may replay a captured graph while + # others run eagerly, causing EP collectives to hang. + if self.has_mtp_cuda_graphs: + unwrapped_model.use_mtp_cuda_graphs = context.using_cuda_graph_this_step() + # If using symmetric kernels and we are using using nccl # for prefill turn off symmetric kernels symmetric_ar_type = self.model_config.symmetric_ar_type @@ -697,118 +746,77 @@ def _dynamic_step_sample_bookkeeping(self): bucket_map[sampling_params].append(request_index) # Just unpack the key directly! + device = torch.cuda.current_device() self._torch_sampling_buckets = [ (indices, *sampling_params) for sampling_params, indices in bucket_map.items() ] + # Pre-compute index tensors on GPU to avoid per-step H2D copies. + self._torch_sampling_bucket_index_tensors = [ + torch.tensor(indices, device=device, dtype=torch.long) + for indices, *_ in self._torch_sampling_buckets + ] - def _rewind_kv_cache(self): + def _rewind_kv_cache(self) -> tuple: """Update the KV cache bookkeeping for speculative decoding. After forward pass with speculative tokens, some tokens may be rejected. - This function "rewinds" the KV cache bookkeeping to reflect only the accepted tokens. - - When speculative tokens are rejected, we need to: - 1. Update request_kv_length_offsets (total sequence length) - 2. Update request_last_kv_block_offset (position within last block) - 3. If rewinding crosses a block boundary: - - Reduce request_kv_block_counts - - Update request_last_kv_block_id to point to the previous block - - Clear the entry in request_to_kv_block_ids for the released block - - Release the block back to the allocator + This function "rewinds" the KV cache bookkeeping to reflect only the accepted + tokens. The core bookkeeping is handled by a Triton kernel (one thread per + request). Mamba hybrid-model state updates remain in PyTorch. + + Returns (blocks_to_release, remove_mask) for the caller to release blocks + back to the allocator outside the compiled graph. """ 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) - # Get the accepted token counts for each request - # Note: _accepted_token_counts is indexed from 0 to active_request_count-1 accepted_tokens_per_request = self._accepted_token_counts_per_request[:active_request_count] - # Number of tokens to rewind (rejected speculative tokens) - num_tokens_to_rewind = self.num_speculative_tokens - accepted_tokens_per_request - - # For prefill requests, no speculative tokens were forwarded through the model, - # so there is nothing to rewind. request_in_prefill_status = context.request_in_prefill_status_tensor[active_request_slice] - num_tokens_to_rewind[request_in_prefill_status == 1] = 0 - - # Save the original offset BEFORE modifying to correctly detect block boundary crossing - original_offset = context.request_last_kv_block_offset[active_request_slice].clone() - - # Check which requests need to rewind to a previous block BEFORE modifying - # A request crosses back to a previous block if: original_offset - num_tokens_to_rewind < 0 - remove_allocated_blocks_mask = (original_offset - num_tokens_to_rewind) < 0 - - # Update the offsets - context.request_last_kv_block_offset[active_request_slice] = ( - original_offset - num_tokens_to_rewind - ) % context.block_size_tokens - - context.request_kv_length_offsets[active_request_slice] = ( - context.request_kv_length_offsets[active_request_slice] - num_tokens_to_rewind + request_last_kv_block_offset = context.request_last_kv_block_offset[active_request_slice] + request_kv_length_offsets = context.request_kv_length_offsets[active_request_slice] + request_kv_block_counts = context.request_kv_block_counts[active_request_slice] + request_last_kv_block_id = context.request_last_kv_block_id[active_request_slice] + request_to_kv_block_ids = context.request_to_kv_block_ids[active_request_slice] + + # --- Triton kernel: core KV-cache rewind --- + blocks_to_release, remove_mask = rewind_kv_cache( + accepted_counts=accepted_tokens_per_request, + prefill_status=request_in_prefill_status, + last_kv_block_offset=request_last_kv_block_offset, + kv_length_offsets=request_kv_length_offsets, + kv_block_counts=request_kv_block_counts, + last_kv_block_id=request_last_kv_block_id, + kv_block_ids=request_to_kv_block_ids, + num_speculative_tokens=self.num_speculative_tokens, + block_size_tokens=context.block_size_tokens, + num_active_requests=active_request_count, ) - # No need to update request_query_lengths (It will be set correctly in the next iteration) - - # For requests that crossed back to a previous block, we need to: - # 1. Reduce the block count by 1 - # 2. Get the block ID to release (current request_last_kv_block_id) - # 3. Update request_last_kv_block_id to point to the previous block - # 4. Clear the entry in request_to_kv_block_ids for the released block - # 5. Release the block back to the allocator - if remove_allocated_blocks_mask.any(): - # Get indices of requests that need to release a block (relative to active requests) - requests_needing_release = torch.nonzero(remove_allocated_blocks_mask, as_tuple=True)[0] - # Convert to absolute indices in the context tensors - absolute_indices = requests_needing_release + context.paused_request_count - - # No clone needed: advanced (fancy) indexing with a tensor already returns - # a copy, not a view. - blocks_to_release = context.request_last_kv_block_id[absolute_indices] - - # Reduce block counts for requests that crossed back - context.request_kv_block_counts[absolute_indices] -= 1 - - # Get the new block counts after decrement - new_block_counts = context.request_kv_block_counts[absolute_indices] - - # Update request_last_kv_block_id to point to the previous block - # and clear the released block entry in request_to_kv_block_ids - # Vectorized implementation using advanced indexing: - # Note: new_block_counts is guaranteed to be > 0 for all requests here, since - # crossing back to a previous block implies the request had at least 2 blocks. - - # Update request_last_kv_block_id to point to the previous block (at index new_count - 1) - context.request_last_kv_block_id[absolute_indices] = context.request_to_kv_block_ids[ - absolute_indices, new_block_counts - 1 - ] - - # Clear the released block entry (at index new_count, which was the old last block) - context.request_to_kv_block_ids[absolute_indices, new_block_counts] = -1 - - # Release the blocks back to the allocator - context.kv_block_allocator.release_memory_blocks(blocks_to_release) - - # Mamba speculative rewind state update + # Mamba speculative rewind: copy accepted intermediate states in-place. if context.is_hybrid_model: - active_mamba_indices = context.mamba_metadata.request_to_mamba_state_idx[ + mamba_state_idx = context.mamba_metadata.request_to_mamba_state_idx[ active_request_slice ] - is_decode_mask = context.request_in_prefill_status_tensor[active_request_slice] == 0 - decode_mamba_indices = active_mamba_indices[is_decode_mask] - accepted_tokens_per_decode_request = accepted_tokens_per_request[is_decode_mask] - - if decode_mamba_indices.numel() > 0: - context.mamba_conv_states[:, decode_mamba_indices] = ( - context.mamba_intermediate_conv_states[ - :, decode_mamba_indices, accepted_tokens_per_decode_request - ] - ) - context.mamba_ssm_states[:, decode_mamba_indices] = ( - context.mamba_intermediate_ssm_states[ - :, decode_mamba_indices, accepted_tokens_per_decode_request - ] - ) + mamba_state_selective_copy( + intermediate_states=context.mamba_intermediate_conv_states, + current_states=context.mamba_conv_states, + prefill_status=request_in_prefill_status, + state_idx=mamba_state_idx, + accepted_counts=accepted_tokens_per_request, + num_layers=context.num_mamba_layers, + ) + mamba_state_selective_copy( + intermediate_states=context.mamba_intermediate_ssm_states, + current_states=context.mamba_ssm_states, + prefill_status=request_in_prefill_status, + state_idx=mamba_state_idx, + accepted_counts=accepted_tokens_per_request, + num_layers=context.num_mamba_layers, + ) + + return blocks_to_release, remove_mask def _sample_from_logits_2d(self, logits_2d: Tensor) -> Tensor: """Sample tokens from 2D logits using existing sampling parameters. @@ -820,18 +828,15 @@ def _sample_from_logits_2d(self, logits_2d: Tensor) -> Tensor: Tensor: Sampled tokens of shape [num_requests]. """ spec_token_list = [] - indices_list = [] - for request_indices, temp, top_k, top_p in self._torch_sampling_buckets: - request_indices_tensor = torch.tensor( - request_indices, device=logits_2d.device, dtype=torch.long - ) + for idx_tensor, (_, temp, top_k, top_p) in zip( + self._torch_sampling_bucket_index_tensors, self._torch_sampling_buckets + ): spec_token_list.append( - self._torch_sampling_func(logits_2d[request_indices_tensor, :], temp, top_k, top_p) + self._torch_sampling_func(logits_2d[idx_tensor, :], temp, top_k, top_p) ) - indices_list.append(request_indices_tensor) spec_tokens = torch.empty(logits_2d.shape[0], device=logits_2d.device, dtype=torch.int64) - for tokens, indices in zip(spec_token_list, indices_list): + for tokens, indices in zip(spec_token_list, self._torch_sampling_bucket_index_tensors): spec_tokens[indices] = tokens return spec_tokens @@ -847,14 +852,15 @@ def _compute_serial_mtp_and_sample(self): (scattered along the first dimension) between MTP depths to avoid a redundant gather + scatter round-trip per depth. """ + nvtx_range_push("mtp-spec-decoding/serial-mtp-init") context = self.inference_wrapped_model.inference_context active_request_count = context.total_request_count - context.paused_request_count active_slice = slice(context.paused_request_count, context.total_request_count) - unwrapped_model = unwrap_model(self.inference_wrapped_model.model) + unwrapped_model = self._unwrapped_model # On non-last pipeline stages, the model won't have decoder hidden states. - has_mtp = is_pipeline_last_stage(self.pp_group) and hasattr( + has_mtp = self._is_last_pp_stage and hasattr( unwrapped_model, '_decoder_hidden_states_cache' ) @@ -865,7 +871,7 @@ def _compute_serial_mtp_and_sample(self): # When SP is active the decoder output is in scattered format # [S/TP, B, H], but _last_accepted_seq_indices are indices into # the full (gathered) sequence. - if self.model_config.sequence_parallel: + if self._sp_enabled: hidden_states = gather_from_sequence_parallel_region( hidden_states, group=self.inference_wrapped_model.tp_group ) @@ -880,72 +886,88 @@ def _compute_serial_mtp_and_sample(self): # The next position to predict starts at that cache length. adjusted_offsets = context.request_kv_length_offsets[active_slice] processed_tokens = context.request_query_lengths[active_slice] - base_position = adjusted_offsets + processed_tokens + # Cast to int64 to match CUDA graph capture dtype expectations. + base_position = (adjusted_offsets + processed_tokens).to(torch.int64) # Start with the freshly sampled base token. next_token_ids = self._sampled_tokens_cuda[:active_request_count].clone() current_hidden = last_accepted_hidden if has_mtp else None - # Compute padding needed to make batch a multiple of tp_size for SP compatibility. - tp_size = get_pg_size(self.inference_wrapped_model.tp_group) - sp_enabled = self.model_config.sequence_parallel and tp_size > 1 - if sp_enabled: - pad_count = (tp_size - active_request_count % tp_size) % tp_size - padded_count = active_request_count + pad_count + # Compute padding needed to make batch compatible with SP and CUDA graphs. + if getattr(self, '_mtp_resolved_padded_count', None) is not None: + # CUDA-graph path: use the EP-synced padded count. + padded_count = self._mtp_resolved_padded_count + assert not self._sp_enabled or padded_count % self._tp_size == 0 + elif has_mtp: + # Eager path: pad only for SP alignment. + padded_count = active_request_count + if self._sp_enabled: + padded_count = round_up_to_nearest_multiple(padded_count, self._tp_size) else: - pad_count = 0 + padded_count = active_request_count + pad_count = padded_count - active_request_count + + # Pad hidden states and scatter for sequence parallelism. + if has_mtp: + current_hidden = F.pad(current_hidden, (0, 0, 0, 0, 0, pad_count)) + if self._sp_enabled: + current_hidden = scatter_to_sequence_parallel_region( + current_hidden, group=self.inference_wrapped_model.tp_group + ) - # Pad hidden states to align with the tensor parallel size. - if has_mtp and sp_enabled: - if pad_count > 0: - current_hidden = F.pad(current_hidden, (0, 0, 0, 0, 0, pad_count)) + token_ids_buf = self._mtp_token_ids_buf[:, :padded_count] + position_ids_buf = self._mtp_position_ids_buf[:, :padded_count] - current_hidden = scatter_to_sequence_parallel_region( - current_hidden, group=self.inference_wrapped_model.tp_group - ) + # Zero-fill padding slots so the embedding layer never sees out-of-range IDs. + token_ids_buf[0, active_request_count:] = 0 + position_ids_buf[0, active_request_count:] = 0 - num_depths = min(self.num_speculative_tokens, self.num_mtp_heads) - for depth in range(num_depths): - position_ids = (base_position + depth).unsqueeze(0) # [1, active_request_count] - token_ids = next_token_ids.unsqueeze(0) # [1, active_request_count] + nvtx_range_pop("mtp-spec-decoding/serial-mtp-init") + for depth in range(self._num_mtp_depths): + nvtx_range_push(f"mtp-spec-decoding/depth-{depth}") + + token_ids_buf[0, :active_request_count] = next_token_ids + position_ids_buf[0, :active_request_count] = base_position + depth mtp_logits_2d = None if has_mtp: - # Pad token_ids and position_ids each iteration (they change per depth). - if pad_count > 0: - token_ids = F.pad(token_ids, (0, pad_count)) - position_ids = F.pad(position_ids, (0, pad_count)) - + nvtx_range_push(f"mtp-spec-decoding/depth-{depth}/forward") + mtp_depth = None if unwrapped_model.mtp.mtp_use_repeated_layer else depth current_hidden, mtp_logits = unwrapped_model.compute_mtp_single_step( hidden_states=current_hidden, - next_token_ids=token_ids, - position_ids=position_ids, - depth=depth, + next_token_ids=token_ids_buf, + position_ids=position_ids_buf, + depth=mtp_depth, ) + nvtx_range_pop(f"mtp-spec-decoding/depth-{depth}/forward") - # Strip padding from logits only. Hidden states stay padded+SP + # Strip padding from logits only. Hidden states stay padded+SP # between depths to avoid redundant gather/scatter round-trips. - if pad_count > 0: - mtp_logits = mtp_logits[:active_request_count] + mtp_logits = mtp_logits[:active_request_count] # mtp_logits: [active_request_count, 1, vocab_size] mtp_logits_2d = mtp_logits.squeeze(1) # [active_request_count, vocab_size] # Broadcast MTP logits across pipeline stages. if self.model_is_pipeline_parallel: + nvtx_range_push(f"mtp-spec-decoding/depth-{depth}/pp-broadcast") mtp_logits_2d = broadcast_from_last_pipeline_stage( [active_request_count, self.vocab_size], dtype=self.model_config.params_dtype, tensor=mtp_logits_2d, pp_group=self.pp_group, ) + nvtx_range_pop(f"mtp-spec-decoding/depth-{depth}/pp-broadcast") # Sample speculative token using the same sampling parameters. + nvtx_range_push(f"mtp-spec-decoding/depth-{depth}/sample") spec_tokens = self._sample_from_logits_2d(mtp_logits_2d) self._sampled_mtp_tokens_cuda[depth, :active_request_count] = spec_tokens + nvtx_range_pop(f"mtp-spec-decoding/depth-{depth}/sample") # Use sampled token as input for the next depth. next_token_ids = spec_tokens + nvtx_range_pop(f"mtp-spec-decoding/depth-{depth}") # Clean up cached hidden states. if has_mtp: @@ -984,13 +1006,10 @@ def _sample_speculative_logits( output_tokens_jumbled_list = [] token_order_list = [] - for request_indices, temp, top_k, top_p in self._torch_sampling_buckets: - request_indices_tensor = torch.tensor( - request_indices, device=token_to_request_index.device - ) - required_indices = torch.where( - torch.isin(token_to_request_index, request_indices_tensor) - )[0] + for idx_tensor, (_, temp, top_k, top_p) in zip( + self._torch_sampling_bucket_index_tensors, self._torch_sampling_buckets + ): + required_indices = torch.where(torch.isin(token_to_request_index, idx_tensor))[0] output_tokens_jumbled_list.append( self._torch_sampling_func(required_logits[required_indices, :], temp, top_k, top_p) ) @@ -1012,86 +1031,19 @@ def _verify_speculative_tokens( self, output_tokens: Tensor, input_tokens_required: Tensor, - request_in_prefill_status_tensor: Tensor, - repeats: Tensor, num_decode_requests: int, num_prefill_requests: int, active_request_count: int, ) -> tuple: - """Verify speculative tokens against input tokens and compute acceptance. - - Creates an accepted tokens mask where: - - For prefill requests, the token is always accepted. - - For decode requests, the first token (base token) is always accepted, then we compare - sampled tokens with input tokens and accept consecutive matches. - Then finds the index of the last accepted token per request. - - Example (assume 1, 2, and 0 spec tokens are accepted in the first 3 decode requests): - input_tokens_required: [ a5 a6s a7s | b3 b4s b5s | c6 c7s c8s | d2 | e4 ] # Size 11 - Output tokens [ a6o a7o a8o | b40 b5o b6o | c7o c8o c9o | d3o | e5o ] - Output tokens right shift [ d3o a6o a7o | a8o b40 b5o | b6o c7o c8o | c9o | d3o ] - Accepted tokens mask [ 1 1 0 | 1 1 1 | 1 0 0 | 1 | 1 ] - Last one indices [ 1 | 5 | 6 | 9 | 10 ] - - Returns: - tuple: (last_one_indices, accepted_tokens_mask, input_tokens_required) where - last_one_indices contains the index of the last accepted token per request. - """ - if input_tokens_required.ndim == 2: - assert ( - input_tokens_required.shape[0] == 1 - ), f"Expected input_tokens_required to have 1 row, but got {input_tokens_required.shape}" - input_tokens_required = input_tokens_required.squeeze(0) - - # Initialize mask with False to prevent boundary bleed - accepted_tokens_mask = torch.zeros_like(input_tokens_required, dtype=torch.bool) - - # Make all prefill tokens accepted - token_to_prefill_idx = torch.repeat_interleave(request_in_prefill_status_tensor, repeats) - accepted_tokens_mask[token_to_prefill_idx == 1] = True - - # Safe decode token verification without cross-batch boundary contamination - decode_mask_2d = None - if num_decode_requests > 0: - decode_len = num_decode_requests * (self.num_speculative_tokens + 1) - - decode_inputs = input_tokens_required[:decode_len].reshape( - num_decode_requests, self.num_speculative_tokens + 1 - ) - decode_outputs = output_tokens[:decode_len].reshape( - num_decode_requests, self.num_speculative_tokens + 1 - ) - - # Shift outputs right by 1 *within* each request to align sampled tokens with input targets - decode_outputs_shifted = decode_outputs.roll(1, dims=1) - decode_mask_2d = decode_inputs == decode_outputs_shifted - # The first token (base token) is always accepted - decode_mask_2d[:, 0] = True - # Enforce consecutive acceptance: cummin propagates False to the right - decode_mask_2d = decode_mask_2d.cummin(dim=1).values - accepted_tokens_mask[:decode_len] = decode_mask_2d.flatten() - - last_one_indices = torch.full( - (active_request_count,), -1, device=input_tokens_required.device + """Verify speculative tokens against input tokens (Triton kernel).""" + return verify_speculative_tokens( + input_tokens=input_tokens_required, + output_tokens=output_tokens, + num_decode_requests=num_decode_requests, + num_prefill_requests=num_prefill_requests, + num_speculative_tokens=self.num_speculative_tokens, ) - if num_decode_requests > 0: - # Summing the consecutive mask gives the count; subtract 1 for the local index - local_last_indices = decode_mask_2d.sum(dim=1) - 1 - row_offsets = torch.arange(num_decode_requests, device=last_one_indices.device) * ( - self.num_speculative_tokens + 1 - ) - last_one_indices[:num_decode_requests] = row_offsets + local_last_indices - - if num_prefill_requests > 0: - decode_len = num_decode_requests * (self.num_speculative_tokens + 1) - prefill_valid = ( - torch.nonzero(accepted_tokens_mask[decode_len:]).squeeze(-1) + decode_len - ) - last_one_indices[num_decode_requests:] = prefill_valid - - return last_one_indices, accepted_tokens_mask, input_tokens_required - def _dynamic_step_sample_logits_and_verify_tokens(self, logits: Tensor, input_ids: Tensor): """ Sample tokens from logits for dynamic batching with speculative tokens and verify the tokens. @@ -1102,16 +1054,11 @@ def _dynamic_step_sample_logits_and_verify_tokens(self, logits: Tensor, input_id request_in_prefill_status_tensor = context.request_in_prefill_status_tensor[ context.paused_request_count : context.total_request_count ] - request_query_lengths = context.request_query_lengths[ - context.paused_request_count : context.total_request_count - ] - - num_prefill_requests = request_in_prefill_status_tensor.sum().item() - num_decode_requests = active_request_count - num_prefill_requests # Get the logit indices for tokens that need sampling. # These indices are always needed for input_ids slicing and tracking # accepted sequence positions, even when logits are pre-sliced. + nvtx_range_push("mtp-spec-decoding/verify/logit-indices") required_logit_indices = context.speculative_required_logit_indices(logits.device) if context.config.materialize_only_last_token_logits: @@ -1121,57 +1068,76 @@ def _dynamic_step_sample_logits_and_verify_tokens(self, logits: Tensor, input_id required_logits = logits.squeeze(0)[ required_logit_indices, : ] # Shape [num_required, vocab_size] + nvtx_range_pop("mtp-spec-decoding/verify/logit-indices") # Sample tokens from logits + nvtx_range_push("mtp-spec-decoding/verify/sample") output_tokens, repeats = self._sample_speculative_logits( required_logits, request_in_prefill_status_tensor ) + nvtx_range_pop("mtp-spec-decoding/verify/sample") + + num_prefill_requests = context.num_prefill_requests + num_decode_requests = active_request_count - num_prefill_requests # Verify speculative tokens against input tokens. + nvtx_range_push("mtp-spec-decoding/verify/verify-tokens") input_tokens_required = input_ids[0, required_logit_indices] last_one_indices, accepted_tokens_mask, input_tokens_required = ( self._verify_speculative_tokens( output_tokens, input_tokens_required, - request_in_prefill_status_tensor, - repeats, num_decode_requests, num_prefill_requests, active_request_count, ) ) + nvtx_range_pop("mtp-spec-decoding/verify/verify-tokens") + + nvtx_range_push("mtp-spec-decoding/verify/prepare-next") + self._prepare_speculative_tokens_for_next_forward_pass( + num_decode_requests, + output_tokens, + required_logit_indices, + last_one_indices, + accepted_tokens_mask, + input_tokens_required, + ) + nvtx_range_pop("mtp-spec-decoding/verify/prepare-next") - # Store the final sampled tokens for the next forward pass. - final_sampled_tokens = output_tokens[last_one_indices] - self._sampled_tokens_cuda[: len(final_sampled_tokens)] = final_sampled_tokens - - # Store the last accepted positions in the packed sequence for serial - # MTP computation after verification. - self._last_accepted_seq_indices = required_logit_indices[last_one_indices] - - # Extract accepted tokens and counts for decode requests. - # For prefill it is always set to 1. For decode, the first token is always accepted, - # then we compare with input tokens and accept the next tokens if its a match. - # - # Example (continuing from above): - # input_tokens_required: [ a5 a6s a7s | b3 b4s b5s | c6 c7s c8s | d2 | e4 ] - # Accepted tokens mask [ 1 1 0 | 1 1 1 | 1 0 0 | 1 | 1 ] - # Accepted tokens [ [a6s -1] | [b4s b5s] | [-1 -1] ] # Only decode requests (prefill defaults to -1) - # Accepted token counts [ 1 | 2 | 0 ] # Prefill defaults to 0 - input_tokens_required[accepted_tokens_mask == 0] = -1 # Mask out non-accepted tokens - input_tokens_decode_mode = input_tokens_required[ - : num_decode_requests * (self.num_speculative_tokens + 1) - ] - input_tokens_reshaped = input_tokens_decode_mode.reshape( - -1, self.num_speculative_tokens + 1 - ) # shape: [num_decode_requests, num_speculative_tokens + 1] - - # Skip the first token of every decode request (i.e a5, b3, c6) - accepted_tokens = input_tokens_reshaped[:, 1:] - self._accepted_tokens_per_request[: accepted_tokens.shape[0], :] = accepted_tokens - self._accepted_token_counts_per_request = (self._accepted_tokens_per_request != -1).sum( - dim=1 + def _prepare_speculative_tokens_for_next_forward_pass( + self, + num_decode_requests: int, + output_tokens: torch.Tensor, + required_logit_indices: torch.Tensor, + last_one_indices: torch.Tensor, + accepted_tokens_mask: torch.Tensor, + input_tokens_required: torch.Tensor, + ): + """Prepare accepted speculative tokens for the next forward pass (Triton kernel). + + Example: + input_tokens_required: [ a5 a6s a7s | b3 b4s b5s | c6 c7s c8s | d2 | e4 ] + Accepted tokens mask [ 1 1 0 | 1 1 1 | 1 0 0 | 1 | 1 ] + Accepted tokens [ [a6s -1] | [b4s b5s] | [-1 -1] ] (decode only; prefill → -1) + Accepted token counts [ 1 | 2 | 0 ] (prefill defaults to 0) + """ + active_request_count = last_one_indices.shape[0] + prepare_next_forward_pass( + num_decode_requests=num_decode_requests, + output_tokens=output_tokens, + required_logit_indices=required_logit_indices, + last_one_indices=last_one_indices, + accepted_tokens_mask=accepted_tokens_mask, + input_tokens=input_tokens_required, + sampled_tokens_buf=self._sampled_tokens_cuda, + last_accepted_seq_buf=self._last_accepted_seq_indices_buf, + accepted_tokens_per_request=self._accepted_tokens_per_request, + accepted_token_counts=self._accepted_token_counts_per_request, + num_speculative_tokens=self.num_speculative_tokens, ) + # Expose the active slice so downstream code sees the right length. + self._last_accepted_seq_indices = self._last_accepted_seq_indices_buf[:active_request_count] def _dynamic_step_sample_logits(self, logits: Tensor): """Sample tokens from logits for dynamic batching. @@ -1628,7 +1594,8 @@ def dummy_forward(self): if not context.cuda_graph_batch_dimensions_list: self.inference_wrapped_model.dummy_forward() - # Disable MoE padding for MTP computation + # Disable MoE padding for MTP computation. + # No CUDA graphs in this path (cuda_graph_batch_dimensions_list is empty). if self.model_config.moe_pad_experts_for_cuda_graph_inference: unwrapped_model = unwrap_model(self.inference_wrapped_model.model) set_decode_expert_padding(unwrapped_model, False) @@ -1651,10 +1618,12 @@ def dummy_forward(self): # fallback to eager dummy forward self.inference_wrapped_model.dummy_forward() - # Disable MoE padding for MTP computation + # Disable MoE padding for MTP computation, unless CUDA graphs + # are active (the graphs were captured with padding enabled). if self.model_config.moe_pad_experts_for_cuda_graph_inference: - unwrapped_model = unwrap_model(self.inference_wrapped_model.model) - set_decode_expert_padding(unwrapped_model, False) + if not context.using_cuda_graph_this_step(): + unwrapped_model = unwrap_model(self.inference_wrapped_model.model) + set_decode_expert_padding(unwrapped_model, False) # When speculative decoding is active, the real EP ranks perform serial # MTP forward passes after the main forward pass. MTP layers may contain @@ -1686,10 +1655,11 @@ def _dummy_serial_mtp_forward(self): if self.model_config.expert_model_parallel_size <= 1: return - unwrapped_model = unwrap_model(self.inference_wrapped_model.model) + unwrapped_model = self._unwrapped_model - is_last_stage = is_pipeline_last_stage(self.pp_group) - has_mtp = is_last_stage and hasattr(unwrapped_model, '_decoder_hidden_states_cache') + has_mtp = self._is_last_pp_stage and hasattr( + unwrapped_model, '_decoder_hidden_states_cache' + ) if not has_mtp and not self.model_is_pipeline_parallel: # No MTP on this rank and no PP broadcast to participate in. return @@ -1697,31 +1667,38 @@ def _dummy_serial_mtp_forward(self): device = torch.cuda.current_device() dtype = self.model_config.params_dtype hidden_size = self.model_config.hidden_size - num_depths = min(self.num_speculative_tokens, self.num_mtp_heads) - # Pad token_ids/position_ids to nearest multiple of tp_size so that the - # embedding can reduce-scatter evenly across TP ranks. - tp_size = get_pg_size(self.inference_wrapped_model.tp_group) - sp_enabled = self.model_config.sequence_parallel and tp_size > 1 - padded_count = tp_size if sp_enabled else 1 + # Use precomputed MTP CUDA graph batch size when available; + # otherwise use minimal SP-compatible size. + if getattr(self, '_mtp_resolved_padded_count', None) is not None: + padded_count = self._mtp_resolved_padded_count + assert not self._sp_enabled or padded_count % self._tp_size == 0 + elif has_mtp: + # Eager path: use TP-aligned minimum size for dummy tensors. + padded_count = self._tp_size if self._sp_enabled else 1 dummy_hidden = None if has_mtp: - # Minimal dummy tensors — just enough to drive the MTP layer forward + # Minimal dummy tensors to drive the MTP layer forward # so that the MoE all-to-all collectives are issued. - # Depth 0 uses full-format hidden; subsequent depths use SP format. - dummy_hidden = torch.zeros((1, 1, hidden_size), device=device, dtype=dtype) + dummy_hidden = torch.zeros((padded_count, 1, hidden_size), device=device, dtype=dtype) + if self._sp_enabled: + dummy_hidden = scatter_to_sequence_parallel_region( + dummy_hidden, group=self.inference_wrapped_model.tp_group + ) dummy_token_ids = torch.zeros((1, padded_count), device=device, dtype=torch.long) dummy_position_ids = torch.zeros((1, padded_count), device=device, dtype=torch.long) - for depth in range(num_depths): + for depth in range(self._num_mtp_depths): + nvtx_range_push(f"mtp-spec-decoding/dummy-depth-{depth}") mtp_logits_2d = None if has_mtp: + mtp_depth = None if unwrapped_model.mtp.mtp_use_repeated_layer else depth dummy_hidden, mtp_logits = unwrapped_model.compute_mtp_single_step( hidden_states=dummy_hidden, next_token_ids=dummy_token_ids, position_ids=dummy_position_ids, - depth=depth, + depth=mtp_depth, ) mtp_logits_2d = mtp_logits.squeeze(1) # [padded_count, vocab_size] @@ -1733,6 +1710,7 @@ def _dummy_serial_mtp_forward(self): tensor=mtp_logits_2d, pp_group=self.pp_group, ) + nvtx_range_pop(f"mtp-spec-decoding/dummy-depth-{depth}") def _dynamic_step_context_bookkeeping(self) -> Dict[str, Tensor]: """Update the dynamic inference context after sampling. @@ -1887,17 +1865,28 @@ async def async_generate_output_tokens_dynamic_batch( if self.num_speculative_tokens > 0: # Phase 1: Verify speculative tokens using base logits only. + nvtx_range_push("mtp-spec-decoding/verify") self._dynamic_step_sample_logits_and_verify_tokens(logits, input_ids) + nvtx_range_pop("mtp-spec-decoding/verify") # Phase 2: Rewind KV cache for rejected tokens. - self._rewind_kv_cache() + nvtx_range_push("mtp-spec-decoding/rewind-kv-cache") + blocks_to_release, remove_mask = self._rewind_kv_cache() + nvtx_range_pop("mtp-spec-decoding/rewind-kv-cache") - # Disable MoE padding for MTP computation + # Disable MoE padding for MTP computation, unless CUDA graphs + # are active (the graphs were captured with padding enabled). if self.model_config.moe_pad_experts_for_cuda_graph_inference: - unwrapped_model = unwrap_model(self.inference_wrapped_model.model) - set_decode_expert_padding(unwrapped_model, False) + if not context.using_cuda_graph_this_step(): + set_decode_expert_padding(self._unwrapped_model, False) # Phase 3: Compute MTP serially with correct (verified) inputs. + nvtx_range_push("mtp-spec-decoding/serial-mtp") self._compute_serial_mtp_and_sample() + nvtx_range_pop("mtp-spec-decoding/serial-mtp") + + # Phase 4: Release freed blocks. Deferred from Phase 2 so the + # data-dependent boolean-mask sync overlaps with MTP GPU work. + context.kv_block_allocator.release_memory_blocks(blocks_to_release[remove_mask]) else: self._dynamic_step_sample_logits(logits) diff --git a/megatron/core/models/common/language_module/language_module.py b/megatron/core/models/common/language_module/language_module.py index e8bb564e759..85870726269 100644 --- a/megatron/core/models/common/language_module/language_module.py +++ b/megatron/core/models/common/language_module/language_module.py @@ -8,6 +8,7 @@ from megatron.core import parallel_state, tensor_parallel from megatron.core.dist_checkpointing.mapping import ShardedStateDict +from megatron.core.transformer.cuda_graphs import CudaGraphManager try: from megatron.core.extensions.transformer_engine import te_parallel_cross_entropy @@ -63,6 +64,20 @@ def __init__( self.vp_stage = None self.vp_size = self.config.virtual_pipeline_model_parallel_size + def _setup_mtp_cuda_graphs(self): + """Wrap `compute_mtp_single_step` with a CudaGraphManager. + + Must be called by subclasses after `self.mtp` is created. + """ + if self.config.cuda_graph_impl == "local": + self._mtp_cudagraph_manager = CudaGraphManager( + self.config, + base_module=self, + function_name="compute_mtp_single_step", + need_backward=False, + is_mtp_inference=True, + ) + def _is_in_embd_group(self): if self.embd_group is None: return False @@ -325,7 +340,11 @@ def shared_embedding_or_output_weight(self) -> Tensor: @torch.inference_mode() def compute_mtp_single_step( - self, hidden_states: Tensor, next_token_ids: Tensor, position_ids: Tensor, depth: int + self, + hidden_states: Tensor, + next_token_ids: Tensor, + position_ids: Tensor, + depth: Optional[int] = None, ) -> tuple: """Compute a single MTP depth for speculative decoding. @@ -336,13 +355,15 @@ def compute_mtp_single_step( hidden_states (Tensor): Hidden states at last accepted positions. next_token_ids (Tensor): Correct next token IDs [1, N]. position_ids (Tensor): Position IDs for the next tokens [1, N]. - depth (int): MTP depth index (0-indexed). + depth (int, optional): MTP depth index. Only needed when + ``mtp_use_repeated_layer`` is False (each depth uses a + distinct layer). Omit for repeated-layer models so that a + single CUDA graph can serve all depths. Returns: tuple: (new_hidden_states, logits [N, 1, vocab_size]). """ - layer_idx = 0 if self.mtp.mtp_use_repeated_layer else depth - + layer_idx = 0 if depth is None else depth mtp_hidden = self.mtp.layers[layer_idx].forward_single_position( hidden_states=hidden_states, next_token_ids=next_token_ids, diff --git a/megatron/core/models/gpt/gpt_model.py b/megatron/core/models/gpt/gpt_model.py index 4fe641bb17b..dedada837b7 100644 --- a/megatron/core/models/gpt/gpt_model.py +++ b/megatron/core/models/gpt/gpt_model.py @@ -223,6 +223,8 @@ def __init__( pg_collection=self.pg_collection, ) + self._setup_mtp_cuda_graphs() + # Output if self.post_process: diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index b575792e81b..4399c6984a7 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -279,6 +279,7 @@ def __init__( mtp_num_depths=self.mtp_num_depths, hybrid_submodules=hybrid_submodules, ) + self._setup_mtp_cuda_graphs() # Output if post_process or self.mtp_process: diff --git a/megatron/core/tensor_parallel/inference_layers.py b/megatron/core/tensor_parallel/inference_layers.py index 80aa754dd50..14ac28fbefa 100644 --- a/megatron/core/tensor_parallel/inference_layers.py +++ b/megatron/core/tensor_parallel/inference_layers.py @@ -20,6 +20,10 @@ from megatron.core.inference.quantization.utils import mm_mxfp8 from megatron.core.inference.symmetric_memory import SymmetricMemoryManager from megatron.core.model_parallel_config import ModelParallelConfig +from megatron.core.tensor_parallel.mappings import ( + gather_from_tensor_model_parallel_region, + reduce_scatter_to_sequence_parallel_region, +) from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import get_tensor_model_parallel_group_if_none @@ -473,3 +477,75 @@ def forward( else: x = self._matmul_reduce_scatter(x) return x, None + + +def inference_all_gather_from_tensor_model_parallel_region( + x: torch.Tensor, tp_group: torch.distributed.ProcessGroup, config: TransformerConfig +) -> torch.Tensor: + """NVLS-optimized all-gather along the last dimension, with NCCL fallback. + + Replaces `gather_from_tensor_model_parallel_region` in inference paths + where autograd is not needed and NVLS symmetric-memory is available. + + The NVLS path performs a flat all-gather into symmetric memory (concatenating + along dim-0), then rearranges the result to the last dimension — the same + semantics as `_gather_along_last_dim` but using hardware multicast when + possible. + """ + tp_size = dist.get_world_size(tp_group) + if tp_size == 1: + return x + + triton_nvls_kernels_allowed = not getattr( + config, 'inference_disable_triton_nvls_kernels', False + ) + + if triton_nvls_kernels_allowed and SymmetricMemoryManager.is_initialized("tp"): + ag_buffer_dims = list(x.size()) + ag_buffer_dims[0] *= tp_size + buf = SymmetricMemoryManager.get_buffer("tp", process_group=tp_group) + symm_mem_buffer = buf.maybe_get_tensor(ag_buffer_dims, dtype=x.dtype) + + if are_tensors_nvls_eligible(x) and symm_mem_buffer["handle"] is not None: + multimem_all_gather(symm_mem_buffer["tensor"], x, symm_mem_buffer["handle"]) + tensor_list = symm_mem_buffer["tensor"].chunk(tp_size, dim=0) + return torch.cat(tensor_list, dim=-1).contiguous() + + return gather_from_tensor_model_parallel_region(x, group=tp_group) + + +def inference_reduce_scatter_to_sequence_parallel_region( + x: torch.Tensor, tp_group: torch.distributed.ProcessGroup, config: TransformerConfig +) -> torch.Tensor: + """NVLS-optimized reduce-scatter along the first dimension, with NCCL fallback. + + Replaces `reduce_scatter_to_sequence_parallel_region` in inference paths + where autograd is not needed and NVLS symmetric-memory is available. + """ + # TODO(ksanthanam): Refactor InferenceRowParallelLinear._matmul_reduce_scatter + # to use this function for its non-fused NVLS reduce-scatter path. + tp_size = dist.get_world_size(tp_group) + if tp_size == 1: + return x + + triton_nvls_kernels_allowed = not getattr( + config, 'inference_disable_triton_nvls_kernels', False + ) + + if triton_nvls_kernels_allowed and SymmetricMemoryManager.is_initialized("tp"): + buf = SymmetricMemoryManager.get_buffer("tp", process_group=tp_group) + symm_mem_buffer = buf.maybe_get_tensor(list(x.size()), dtype=x.dtype) + + if ( + x.dtype == torch.bfloat16 + and are_tensors_nvls_eligible(x) + and symm_mem_buffer["handle"] is not None + ): + symm_mem_buffer["tensor"].copy_(x) + output_dims = list(x.size()) + output_dims[0] = x.size(0) // tp_size + output = torch.empty(output_dims, dtype=x.dtype, device=x.device) + multimem_reduce_scatter(output, symm_mem_buffer["tensor"], symm_mem_buffer["handle"]) + return output + + return reduce_scatter_to_sequence_parallel_region(x, group=tp_group) diff --git a/megatron/core/tensor_parallel/layers.py b/megatron/core/tensor_parallel/layers.py index 610700f0a95..4ab2aa0f639 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -235,6 +235,11 @@ def __init__( ) self.num_embeddings_per_partition = self.vocab_end_index - self.vocab_start_index self.deterministic_mode = config.deterministic_mode + self.config = config + + self.use_inference_optimized_reduce_scatter = ( + getattr(config, 'transformer_impl', None) == 'inference_optimized' + ) # Allocate weights and initialize. if config.use_cpu_initialization: @@ -302,9 +307,17 @@ def forward(self, input_): if self.reduce_scatter_embeddings: # Data format change to avoid explicit tranposes : [b s h] --> [s b h]. output_parallel = output_parallel.transpose(0, 1).contiguous() - output = reduce_scatter_to_sequence_parallel_region( - output_parallel, group=self.tp_group - ) + if self.use_inference_optimized_reduce_scatter and not self.training: + # Deferred to avoid circular import: inference_layers → TE → layers. + from .inference_layers import inference_reduce_scatter_to_sequence_parallel_region + + output = inference_reduce_scatter_to_sequence_parallel_region( + output_parallel, self.tp_group, self.config + ) + else: + output = reduce_scatter_to_sequence_parallel_region( + output_parallel, group=self.tp_group + ) else: # Reduce across all the model parallel GPUs. output = reduce_from_tensor_model_parallel_region(output_parallel, group=self.tp_group) @@ -921,6 +934,10 @@ def __init__( else: self.register_parameter("bias", None) + self.use_inference_optimized_all_gather = ( + getattr(config, 'transformer_impl', None) == 'inference_optimized' + ) + self.sequence_parallel = config.sequence_parallel if self.sequence_parallel and world_size <= 1: warnings.warn( @@ -1056,7 +1073,17 @@ def forward( if gather_output: # All-gather across the partitions. - output = gather_from_tensor_model_parallel_region(output_parallel, group=self.tp_group) + if self.use_inference_optimized_all_gather and not self.training: + # Deferred to avoid circular import: inference_layers → TE → layers. + from .inference_layers import inference_all_gather_from_tensor_model_parallel_region + + output = inference_all_gather_from_tensor_model_parallel_region( + output_parallel, self.tp_group, self.config + ) + else: + output = gather_from_tensor_model_parallel_region( + output_parallel, group=self.tp_group + ) else: output = output_parallel output_bias = self.bias if self.skip_bias_add else None diff --git a/megatron/core/transformer/cuda_graphs.py b/megatron/core/transformer/cuda_graphs.py index 866f6676d92..c2ebe655955 100644 --- a/megatron/core/transformer/cuda_graphs.py +++ b/megatron/core/transformer/cuda_graphs.py @@ -8,7 +8,7 @@ import os import time from collections import defaultdict -from contextlib import nullcontext +from contextlib import contextmanager, nullcontext from copy import deepcopy from dataclasses import dataclass, is_dataclass from enum import Enum @@ -98,6 +98,16 @@ def _set_capture_end(): _IS_GRAPH_CAPTURING = False +@contextmanager +def graph_capture(): + """Context manager that brackets a graph-capture region.""" + _set_capture_start() + try: + yield + finally: + _set_capture_end() + + def is_graph_warmup(): """Query if currently warming up for graph capture.""" return _IS_GRAPH_WARMUP @@ -331,6 +341,10 @@ class _CudagraphGlobalRecord: cudagraph_record: list[tuple] = [] cudagraph_inference_record: list[tuple] = [] + # MTP CudaGraphManagers registered at construction time so that + # delete_cuda_graphs() can clear their lookup tables. + mtp_cudagraph_managers: list = [] + """A pool-like data structure to reuse input and output buffers across cudagraph.""" tensor_reuse_pool = TensorReusePool() @@ -506,6 +520,19 @@ def delete_cuda_graphs(): runner.bwd_graph = None runner.mempool = None + # Reset MTP runners (excluded from the global inference record). + for mgr in _CudagraphGlobalRecord.mtp_cudagraph_managers: + for runner in mgr.cudagraph_runners: + runner.cudagraph_created = False + runner.fwd_graph_recorded = False + runner.bwd_graph_recorded = False + runner.fwd_graph = None + runner.bwd_graph = None + runner.mempool = None + mgr.cudagraph_runners.clear() + mgr.custom_cudagraphs_lookup_table.clear() + _CudagraphGlobalRecord.mtp_cudagraph_managers.clear() + # Reset global tracking state _CudagraphGlobalRecord.cudagraph_created = False _CudagraphGlobalRecord.cudagraph_record = [] @@ -1417,6 +1444,7 @@ def __init__( function_name=None, need_backward=True, pg_collection=None, + is_mtp_inference=False, inline_capture=False, num_warmup_steps=None, ): @@ -1426,6 +1454,7 @@ def __init__( Args: config: TransformerConfig object containing CUDA graph settings for memory pooling, graph retention, gradient accumulation, FP8/FP4, and warmup steps. + is_mtp_inference: Whether this manager wraps an MTP inference forward pass. inline_capture: Normally, whether the inline capture path is taken depends on whether `inference_context` is present in the kwargs of the forward call. Setting this argument to True always forces the inline capture path to be taken. @@ -1438,6 +1467,7 @@ def __init__( self.pg_collection = pg_collection rng_tracker = get_cuda_rng_tracker() self.need_backward = need_backward + self.is_mtp_inference = is_mtp_inference if function_name is not None: func = getattr(base_module, function_name) @@ -1483,6 +1513,10 @@ def wrapped_func(*args, eager=False, cache_key=None, **kwargs): self.custom_cudagraphs_lookup_table: dict = defaultdict(lambda: None) self.is_first_microbatch = False + if is_mtp_inference: + # Registered so delete_cuda_graphs() can clear the lookup table. + _CudagraphGlobalRecord.mtp_cudagraph_managers.append(self) + # Without pipeline parallelism, microbatches execute one at a time. # Therefore modules will always execute in the same order, so cudagraphs # can both be reused and share a single mempool. @@ -1588,14 +1622,19 @@ def __call__(self, megatron_module, args, kwargs, cache_key=None): cache_key: Optional hashable key for O(1) runner lookup. If `inference_context` is provided, this gets set to the correct value. """ - is_inference_mode = 'inference_context' in kwargs.keys() and kwargs['inference_context'] + is_inference_mode = ( + 'inference_context' in kwargs.keys() and kwargs['inference_context'] + ) or self.is_mtp_inference if cache_key is None and is_inference_mode: - inference_context = kwargs['inference_context'] - if inference_context.is_static_batching(): - batch_size = kwargs['hidden_states'].shape[0] - cache_key = (batch_size, inference_context.is_decode_only()) - else: - cache_key = inference_context.padded_batch_dimensions + if 'inference_context' in kwargs and kwargs['inference_context']: + inference_context = kwargs['inference_context'] + if inference_context.is_static_batching(): + batch_size = kwargs['hidden_states'].shape[0] + cache_key = (batch_size, inference_context.is_decode_only()) + else: + cache_key = inference_context.padded_batch_dimensions + elif self.is_mtp_inference: + cache_key = ('mtp', kwargs['hidden_states'].shape, kwargs.get('depth')) is_in_checkpoint_fwd = is_checkpointing() if HAVE_TE_GRAPHS: is_in_checkpoint_fwd = is_in_checkpoint_fwd or is_fp8_activation_recompute_enabled() @@ -1613,11 +1652,28 @@ def __call__(self, megatron_module, args, kwargs, cache_key=None): out = runner.replay_graph_capture(self.is_first_microbatch, args, kwargs) else: if is_inference_mode or self._inline_capture: + # MTP must match the main model's eager/graph mode so all EP + # ranks take the same code path. Skip during graph capture. + if ( + self.is_mtp_inference + and not getattr(megatron_module, 'use_mtp_cuda_graphs', False) + and not is_graph_capturing() + ): + return self.func(*args, **kwargs) + # Inference generation mode creates graphs immediately runner = self.get_cudagraph_runner( megatron_module, args, kwargs, True, cache_key=cache_key ) + if ( + not runner.fwd_graph_recorded + and self.is_mtp_inference + and not is_graph_capturing() + ): + # No pre-warmed graph for this batch size — run eagerly. + return self.func(*args, **kwargs) + if not runner.fwd_graph_recorded: # Reuse graph input-output buffers for inference local_args, local_kwargs = args, kwargs @@ -1649,10 +1705,14 @@ def __call__(self, megatron_module, args, kwargs, cache_key=None): runner.cudagraph_created = True runner = runner.eval() - # Record this to the global execution record - _CudagraphGlobalRecord.cudagraph_inference_record.append( - (runner, "fwd", args, kwargs) - ) + # Record to the global execution record. MTP runners are + # excluded — they don't chain with decoder layers (the + # previous-layer lookup expects layer_number) and are + # cleaned up via mtp_cudagraph_managers instead. + if not self.is_mtp_inference: + _CudagraphGlobalRecord.cudagraph_inference_record.append( + (runner, "fwd", args, kwargs) + ) # Now replay the graph out = runner.replay_graph_capture(self.is_first_microbatch, args, kwargs) diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index b70fdc3f28b..2e0461e365c 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -22,6 +22,9 @@ gather_from_tensor_model_parallel_region, scatter_to_sequence_parallel_region, ) +from megatron.core.tensor_parallel.inference_layers import ( + inference_all_gather_from_tensor_model_parallel_region, +) from megatron.core.transformer.enums import AttnMaskType, LayerType from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.spec_utils import ModuleSpec, build_module @@ -918,11 +921,15 @@ def _concat_embeddings(self, hidden_states: torch.Tensor, decoder_input: torch.T hidden_states = torch.cat((decoder_input, hidden_states), -1) hidden_states, _ = self.eh_proj(hidden_states) # For tensor parallel we need to gather the tensor across the model-parallel - # ranks after the linear projection. This used to call - # `all_gather_last_dim_from_tensor_parallel_region`, but that utility reduces - # the gradient in backward pass and was therefore incorrect in this context. - # It has been replaced with the correct `gather_from_tensor_model_parallel_region`. - hidden_states = gather_from_tensor_model_parallel_region(hidden_states, group=self.tp_group) + # ranks after the linear projection. + if not self.training: + hidden_states = inference_all_gather_from_tensor_model_parallel_region( + hidden_states, self.tp_group, self.config + ) + else: + hidden_states = gather_from_tensor_model_parallel_region( + hidden_states, group=self.tp_group + ) # For sequence parallel, scatter after linear_fc and before transformer layer. if self.sequence_parallel: hidden_states = scatter_to_sequence_parallel_region(hidden_states, group=self.tp_group) @@ -1021,7 +1028,6 @@ def forward_single_position( rotary_pos_emb: Optional[Tensor] = None, rotary_pos_cos: Optional[Tensor] = None, rotary_pos_sin: Optional[Tensor] = None, - inference_params=None, packed_seq_params: Optional[PackedSeqParams] = None, sequence_len_offset: Optional[Tensor] = None, ) -> Tensor: @@ -1052,7 +1058,6 @@ def forward_single_position( rotary_pos_emb=rotary_pos_emb, rotary_pos_cos=rotary_pos_cos, rotary_pos_sin=rotary_pos_sin, - inference_params=inference_params, packed_seq_params=packed_seq_params, sequence_len_offset=sequence_len_offset, ) diff --git a/megatron/core/utils.py b/megatron/core/utils.py index 39ebf9a5044..f0d00cd40bb 100644 --- a/megatron/core/utils.py +++ b/megatron/core/utils.py @@ -506,6 +506,11 @@ def divide(numerator, denominator): return numerator // denominator +def round_up_to_nearest_multiple(value: int, multiple: int) -> int: + """Round *value* up to the nearest multiple of *multiple*.""" + return math.ceil(value / multiple) * multiple + + def get_tensor_model_parallel_group_if_none(tp_group, is_expert=False, check_initialized=True): """Issue a deprecation warning if tp_group is None and return the default tp group.""" # TODO(zijiey): remove this function later. diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine.py b/tests/unit_tests/inference/engines/test_dynamic_engine.py index e39a58736a5..4f02369d0ed 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine.py @@ -50,11 +50,11 @@ from megatron.core.models.hybrid.hybrid_model import HybridModel from megatron.core.ssm.mamba_mixer import _check_mamba_sequence_packing_support 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.cuda_graphs import delete_cuda_graphs from megatron.core.transformer.enums import CudaGraphScope from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import is_fa_min_version, is_te_min_version -from tests.unit_tests.test_utilities import Utils +from tests.unit_tests.test_utilities import Utils, clear_nvte_env_vars try: from torch_memory_saver import torch_memory_saver # noqa: F401 @@ -178,7 +178,7 @@ class DynamicEngineTestEnv: ) -class TestDynamicInferenceEngine: +class DynamicInferenceEngineTestBase: @classmethod def _build_requests(cls, test_config: DynamicEngineTestConfig) -> List[DynamicInferenceRequest]: @@ -281,11 +281,7 @@ def _build_inference_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, - pipeline_model_parallel_size=test_config.pipeline_model_parallel_size, - ) - + clear_nvte_env_vars() set_rounder(4) # Random state. @@ -398,12 +394,18 @@ def _build_test_env(cls, test_config): ), sequence_parallel=test_config.sequence_parallel, pipeline_dtype=torch.bfloat16, - add_bias_linear=test_config.expert_model_parallel_size == 1, + add_bias_linear=test_config.expert_model_parallel_size == 1 + and not (test_config.transformer_impl == "inference_optimized"), 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, transformer_impl=test_config.transformer_impl, + normalization=( + "RMSNorm" + if test_config.transformer_impl == "inference_optimized" + else "LayerNorm" + ), is_hybrid_model=True, # Needs to be set for correct out_proj init ) @@ -459,10 +461,7 @@ def _build_test_env(cls, test_config): ), ) - # Reset global cuda graph state. - _CudagraphGlobalRecord.cudagraph_created = False - _CudagraphGlobalRecord.cudagraph_record = [] - CudaGraphManager.global_mempool = None + delete_cuda_graphs() # Inference engine. engine = DynamicInferenceEngine(text_generation_controller, inference_context) @@ -565,8 +564,21 @@ def _run_test(cls, **test_config_kwargs): return env + +class TestDynamicInferenceEngine(DynamicInferenceEngineTestBase): + + @classmethod + def setup_class(cls): + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + expert_model_parallel_size=1, + expert_tensor_parallel_size=1, + ) + @classmethod def teardown_class(cls): + delete_cuda_graphs() set_rounder(64) Utils.destroy_model_parallel() @@ -1086,88 +1098,6 @@ def test_log_probs_token_correspondence(self): assert not math.isnan(log_prob) and not math.isinf(log_prob) assert -100.0 <= log_prob <= 0.0 - @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("materialize_only_last_token_logits", [False, True]) - @pytest.mark.parametrize("sequence_parallel", [False, True]) - @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", "hybrid"]) - @pytest.mark.parametrize("transformer_impl", ["local", "inference_optimized"]) - @torch.inference_mode() - def test_parallel_inference( - self, - model_provider, - tp_size, - pp_size, - ep_size, - sequence_parallel, - materialize_only_last_token_logits, - transformer_impl, - ): - 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(): - pytest.skip("Distributed not initialized") - world_size = torch.distributed.get_world_size() - min_world_size = tp_size * pp_size * ep_size - if world_size < min_world_size: - pytest.skip(f"Test requires at least {min_world_size} GPUs") - elif tp_size == 1 and sequence_parallel: - 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 transformer_impl == "inference_optimized": - if ep_size > 1: - pytest.skip( - reason="MoE models are not supported with the inference optimized transformer." - ) - if tp_size > 1 and not sequence_parallel: - pytest.skip( - reason=( - "The inference optimized transformer requires sequence parallelism " - "when tp_size > 1." - ) - ) - if model_provider == "hybrid": - pytest.skip( - reason="Mamba model is not supported with the inference optimized transformer." - ) - - 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, - sequence_parallel=sequence_parallel, - materialize_only_last_token_logits=materialize_only_last_token_logits, - transformer_impl=transformer_impl, - ) - - @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("materialize_only_last_token_logits", [False, True]) - def test_sequence_parallel_fp8_inference(self, materialize_only_last_token_logits: bool): - fp8_available, reason_for_no_fp8 = check_fp8_support() - if not fp8_available: - pytest.skip(reason_for_no_fp8) - - self._run_test( - min_prompt_length=19, - max_prompt_length=19, - tensor_model_parallel_size=4, - sequence_parallel=True, - materialize_only_last_token_logits=True, - fp8=True, - ) - @pytest.mark.internal @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" @@ -2364,7 +2294,7 @@ def mock_mtp_forward(*args, **kwargs): ) return base_logits - def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth): + def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth=None): n = hidden_states.size(0) logits = torch.zeros( n, 1, test_config.vocab_size, device=hidden_states.device, dtype=torch.bfloat16 @@ -2487,7 +2417,7 @@ def mock_deterministic_forward(*args, **kwargs): ) return base_logits - def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth): + def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth=None): n = hidden_states.size(0) # Predict next_token_ids + 1 (continuing the ascending sequence) pred_toks = (next_token_ids + 1).clamp(max=test_config.vocab_size - 1) @@ -2571,7 +2501,7 @@ def mock_deterministic_forward(*args, **kwargs): ) return base_logits - def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth): + def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth=None): n = hidden_states.size(0) # Predict next_token_ids + 1 (continuing the ascending sequence) pred_toks = (next_token_ids + 1).clamp(max=test_config.vocab_size - 1) @@ -2656,7 +2586,7 @@ def mock_deterministic_forward(*args, **kwargs): ) return base_logits - def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth): + def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth=None): n = hidden_states.size(0) # Predict next_token_ids + 1 (continuing the ascending sequence) pred_toks = (next_token_ids + 1).clamp(max=test_config.vocab_size - 1) @@ -3010,7 +2940,7 @@ def mock_safe_forward(*args, **kwargs): ) return base_logits - def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth): + def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth=None): n = hidden_states.size(0) logits = torch.zeros( n, 1, test_config.vocab_size, device=hidden_states.device, dtype=torch.bfloat16 @@ -3223,7 +3153,7 @@ def mock_deterministic_forward(*args, **kwargs): ) return base_logits - def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth): + def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth=None): n = hidden_states.size(0) logits = torch.randn( n, 1, test_config.vocab_size, device=hidden_states.device, dtype=torch.bfloat16 @@ -3343,7 +3273,7 @@ def mock_deterministic_forward(*args, **kwargs): ) return base_logits - def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth): + def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth=None): n = hidden_states.size(0) logits = torch.randn( n, 1, test_config.vocab_size, device=hidden_states.device, dtype=torch.bfloat16 @@ -3473,7 +3403,7 @@ def mock_deterministic_forward(*args, **kwargs): ) return base_logits - def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth): + def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth=None): n = hidden_states.size(0) logits = torch.randn( n, 1, test_config.vocab_size, device=hidden_states.device, dtype=torch.bfloat16 @@ -3822,7 +3752,7 @@ def mock_deterministic_forward(*args, **kwargs): ) return base_logits - def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth): + def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth=None): n = hidden_states.size(0) pred_toks = (next_token_ids + 1).clamp(max=test_config.vocab_size - 1) logits = torch.zeros( @@ -4260,40 +4190,6 @@ def deterministic_mtp(hidden_states, next_token_ids, position_ids, depth): assert isinstance(lp, float) assert -0.1 < lp <= 0.0, f"Token {j}: expected log prob near 0.0, got {lp}" - @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_speculative_decoding_pipeline_parallel(self): - """Test speculative decoding with pipeline parallelism (pp_size=2). - - Verifies that MTP logit broadcasts across pipeline stages don't hang - or produce incorrect results. Each PP stage must participate in the - same number of MTP broadcast rounds. - """ - if not torch.distributed.is_initialized(): - pytest.skip("Distributed not initialized") - world_size = torch.distributed.get_world_size() - pp_size = 2 - if world_size < pp_size: - pytest.skip(f"Test requires at least {pp_size} GPUs") - - env = self._run_test( - model_provider="gpt", - pipeline_model_parallel_size=pp_size, - num_speculative_tokens=2, - num_tokens_to_generate=6, - materialize_only_last_token_logits=False, - ) - - for request in env.requests: - assert ( - request.status == Status.COMPLETED - ), f"Request {request.request_id}: status={request.status}" - num_expected = request.sampling_params.num_tokens_to_generate - assert len(request.generated_tokens) <= num_expected - @pytest.mark.internal @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" @@ -4413,6 +4309,137 @@ def mtp_with_rejection(hidden_states, next_token_ids, position_ids, depth): assert env.engine.context.total_request_count == 0 +class TestDynamicInferenceEngineParallel(DynamicInferenceEngineTestBase): + """Tests that require non-default parallel configs (tp>1, pp>1, or ep>1). + + Each test initializes its own parallel state and tears it down afterward, + so these are separated from TestDynamicInferenceEngine to avoid accumulating + NCCL communicator memory from repeated init/destroy cycles. + """ + + def teardown_method(self, method): + delete_cuda_graphs() + Utils.destroy_model_parallel() + + @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, + pipeline_model_parallel_size=test_config.pipeline_model_parallel_size, + expert_model_parallel_size=test_config.expert_model_parallel_size, + expert_tensor_parallel_size=1, + ) + return super()._build_test_env(test_config) + + @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("materialize_only_last_token_logits", [False, True]) + @pytest.mark.parametrize("sequence_parallel", [False, True]) + @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", "hybrid"]) + @pytest.mark.parametrize("transformer_impl", ["local", "inference_optimized"]) + @torch.inference_mode() + def test_parallel_inference( + self, + model_provider, + tp_size, + pp_size, + ep_size, + sequence_parallel, + materialize_only_last_token_logits, + transformer_impl, + ): + 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(): + pytest.skip("Distributed not initialized") + world_size = torch.distributed.get_world_size() + min_world_size = tp_size * pp_size * ep_size + if world_size < min_world_size: + pytest.skip(f"Test requires at least {min_world_size} GPUs") + elif tp_size == 1 and sequence_parallel: + 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 transformer_impl == "inference_optimized": + if ep_size > 1: + pytest.skip( + reason="MoE models are not supported with the inference optimized transformer." + ) + if tp_size > 1 and not sequence_parallel: + pytest.skip( + reason=( + "The inference optimized transformer requires sequence parallelism " + "when tp_size > 1." + ) + ) + + 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, + sequence_parallel=sequence_parallel, + materialize_only_last_token_logits=materialize_only_last_token_logits, + transformer_impl=transformer_impl, + ) + + @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("materialize_only_last_token_logits", [False, True]) + def test_sequence_parallel_fp8_inference(self, materialize_only_last_token_logits: bool): + fp8_available, reason_for_no_fp8 = check_fp8_support() + if not fp8_available: + pytest.skip(reason_for_no_fp8) + + self._run_test( + min_prompt_length=19, + max_prompt_length=19, + tensor_model_parallel_size=4, + sequence_parallel=True, + materialize_only_last_token_logits=True, + 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_speculative_decoding_pipeline_parallel(self): + """Test speculative decoding with pipeline parallelism (pp_size=2).""" + if not torch.distributed.is_initialized(): + pytest.skip("Distributed not initialized") + world_size = torch.distributed.get_world_size() + pp_size = 2 + if world_size < pp_size: + pytest.skip(f"Test requires at least {pp_size} GPUs") + + env = self._run_test( + model_provider="gpt", + pipeline_model_parallel_size=pp_size, + num_speculative_tokens=2, + num_tokens_to_generate=6, + materialize_only_last_token_logits=False, + ) + + for request in env.requests: + assert ( + request.status == Status.COMPLETED + ), f"Request {request.request_id}: status={request.status}" + num_expected = request.sampling_params.num_tokens_to_generate + assert len(request.generated_tokens) <= num_expected + + CHUNKED_CG_BLOCK_SIZE = 256 CHUNKED_CG_VOCAB_SIZE = 10000 CHUNKED_CG_MAX_SEQ_LEN = 2048 @@ -4433,6 +4460,7 @@ def setup_class(cls): @classmethod def teardown_class(cls): + delete_cuda_graphs() set_rounder(64) Utils.destroy_model_parallel() @@ -4497,17 +4525,6 @@ def _create_model(self, model_provider, num_cuda_graphs): model.eval() return model - def _reset_cuda_graph_state(self, model): - """Reset all CUDA graph global and per-module state.""" - _CudagraphGlobalRecord.cudagraph_created = False - _CudagraphGlobalRecord.cudagraph_record = [] - _CudagraphGlobalRecord.cudagraph_inference_record = [] - CudaGraphManager.global_mempool = None - for module in model.modules(): - if isinstance(module, CudaGraphManager): - module.cudagraph_runners.clear() - module.custom_cudagraphs_lookup_table.clear() - def _build_engine(self, model, enable_chunked_prefill, num_cuda_graphs, context_max_tokens): """Build an engine with the given chunked prefill / CUDA graph config.""" set_rounder(4) @@ -4540,7 +4557,7 @@ def _build_engine(self, model, enable_chunked_prefill, num_cuda_graphs, context_ vocab_size=CHUNKED_CG_VOCAB_SIZE, detokenize=lambda tokens: "tokenized_prompt" ), ) - self._reset_cuda_graph_state(model) + delete_cuda_graphs() return DynamicInferenceEngine(controller, context) def _run_to_completion(self, engine, prompts, num_tokens_to_generate): @@ -4575,10 +4592,7 @@ def test_chunked_prefill_cuda_graphs(self, model_provider, chunked_prefill, num_ """Verify generated tokens match across chunked prefill and CUDA graph configs.""" skip_if_mamba_sequence_packing_not_available(model_provider) - # Clear NVTE env vars set by conftest set_env fixture. - os.environ.pop('NVTE_FLASH_ATTN', None) - os.environ.pop('NVTE_FUSED_ATTN', None) - os.environ.pop('NVTE_UNFUSED_ATTN', None) + clear_nvte_env_vars() random.seed(123) torch.manual_seed(123) diff --git a/tests/unit_tests/inference/engines/test_mtp_cuda_graph_inference.py b/tests/unit_tests/inference/engines/test_mtp_cuda_graph_inference.py new file mode 100644 index 00000000000..d6605e88d0d --- /dev/null +++ b/tests/unit_tests/inference/engines/test_mtp_cuda_graph_inference.py @@ -0,0 +1,991 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Tests for CUDA-graphed MTP (Multi-Token Prediction) inference. + +Verifies that: +1. CUDA graph replay produces the same output as eager execution (no extra + padding in the CUDA graphed case). +2. CUDA graphs work correctly with sequence parallelism (padding is applied + to make batch sizes divisible by TP). +3. CUDA graphs work correctly with expert parallelism and dummy ranks. + +Uses DynamicInferenceEngine for CUDA graph warmup so MTP graph capture +logic matches production code exactly. +""" + +import itertools +from unittest import mock + +import pytest +import torch +import torch.distributed as dist + +from megatron.core import parallel_state +from megatron.core.inference.batch_dimensions_utils import InferenceBatchDimensions +from megatron.core.inference.config import InferenceConfig +from megatron.core.inference.contexts.dynamic_context import DynamicInferenceContext +from megatron.core.inference.engines.dynamic_engine import DynamicInferenceEngine +from megatron.core.inference.model_inference_wrappers.gpt.gpt_inference_wrapper import ( + GPTInferenceWrapper, +) +from megatron.core.inference.text_generation_controllers.text_generation_controller import ( + TextGenerationController, +) +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_local_spec, + get_gpt_mtp_block_spec, +) +from megatron.core.models.gpt.gpt_model import GPTModel +from megatron.core.tensor_parallel.mappings import scatter_to_sequence_parallel_region +from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer import TransformerConfig +from megatron.core.transformer.cuda_graphs import _CudagraphGlobalRecord, delete_cuda_graphs +from megatron.core.transformer.enums import AttnBackend +from megatron.core.utils import unwrap_model +from tests.unit_tests.test_utilities import Utils + +# --------------------------------------------------------------------------- # +# TestMTPCudaGraphInference (TP = 2) +# --------------------------------------------------------------------------- # + + +class TestMTPCudaGraphInference: + """Tests for MTP CUDA-graphed inference with tensor parallelism. + + All tests require at least 2 GPUs (TP = 2). Uses DynamicInferenceEngine + for CUDA graph warmup so MTP graph capture matches production code. + """ + + HIDDEN_SIZE = 32 + VOCAB_SIZE = 100 + MAX_SEQ_LEN = 64 + NUM_LAYERS = 4 + NUM_ATTN_HEADS = 4 + TP_SIZE = 2 + + @classmethod + def setup_class(cls): + if Utils.world_size < cls.TP_SIZE: + pytest.skip(f"Need at least {cls.TP_SIZE} GPUs") + Utils.initialize_model_parallel( + tensor_model_parallel_size=cls.TP_SIZE, pipeline_model_parallel_size=1 + ) + + @classmethod + def teardown_class(cls): + delete_cuda_graphs() + Utils.destroy_model_parallel() + + def teardown_method(self): + delete_cuda_graphs() + + # ---- helpers ---------------------------------------------------------- # + + def _build_model( + self, *, sequence_parallel=False, mtp_num_layers=2, mtp_use_repeated_layer=False + ): + """Build a GPT model with MTP layers and local CUDA graph support.""" + model_parallel_cuda_manual_seed(123, inference_rng_tracker=True, force_reset_rng=True) + config = TransformerConfig( + num_layers=self.NUM_LAYERS, + hidden_size=self.HIDDEN_SIZE, + num_attention_heads=self.NUM_ATTN_HEADS, + use_cpu_initialization=True, + attention_backend=AttnBackend.local, + params_dtype=torch.bfloat16, + tensor_model_parallel_size=self.TP_SIZE, + pipeline_model_parallel_size=1, + pipeline_dtype=torch.bfloat16, + mtp_num_layers=mtp_num_layers, + mtp_use_repeated_layer=mtp_use_repeated_layer, + sequence_parallel=sequence_parallel, + cuda_graph_impl="local", + ) + layer_spec = get_gpt_layer_local_spec() + mtp_block_spec = get_gpt_mtp_block_spec( + config=config, spec=layer_spec, use_transformer_engine=False + ) + model = GPTModel( + config=config, + transformer_layer_spec=layer_spec, + vocab_size=self.VOCAB_SIZE, + max_sequence_length=self.MAX_SEQ_LEN, + parallel_output=True, + pre_process=True, + post_process=True, + mtp_block_spec=mtp_block_spec, + ).cuda() + for param in model.parameters(): + param.data = param.data.to(config.params_dtype) + model.eval() + return model + + def _build_engine( + self, + *, + sequence_parallel=False, + mtp_num_layers=2, + mtp_use_repeated_layer=False, + num_speculative_tokens=2, + max_requests=16, + ): + """Build a DynamicInferenceEngine with automatic MTP CUDA graph warmup. + + The engine's `__init__` calls `create_cuda_graphs()` which captures + both decoder and MTP CUDA graphs, matching production warmup exactly. + """ + delete_cuda_graphs() + model = self._build_model( + sequence_parallel=sequence_parallel, + mtp_num_layers=mtp_num_layers, + mtp_use_repeated_layer=mtp_use_repeated_layer, + ) + config = model.config + context = DynamicInferenceContext( + model_config=config, + inference_config=InferenceConfig( + max_sequence_length=self.MAX_SEQ_LEN, + buffer_size_gb=0.5, + materialize_only_last_token_logits=False, + num_speculative_tokens=num_speculative_tokens, + block_size_tokens=256, + max_requests=max_requests, + num_cuda_graphs=-1, + ), + ) + wrapped = GPTInferenceWrapper(model, context) + wrapped.model_is_pipeline_parallel = False + mock_tokenizer = mock.Mock() + ctrl = TextGenerationController(inference_wrapped_model=wrapped, tokenizer=mock_tokenizer) + engine = DynamicInferenceEngine(ctrl, context) + return engine + + @staticmethod + def _get_mtp_warmed_batch_sizes(engine): + """Return the MTP batch sizes (padded req_counts) warmed by the engine. + + These are the `n` values for which MTP CUDA graphs were captured. + Hidden states shape is `[n // tp, 1, H]` with SP, `[n, 1, H]` without. + Token/position IDs are always `[1, n]`. + """ + context = engine.context + model_config = engine.controller.inference_wrapped_model.model.config + tp_size = parallel_state.get_tensor_model_parallel_world_size() + sp_enabled = model_config.sequence_parallel and tp_size > 1 + sizes = set() + for dim in context.cuda_graph_batch_dimensions_list: + n = dim.req_count + if sp_enabled: + n += (tp_size - n % tp_size) % tp_size + if n > 0: + sizes.add(n) + return sorted(sizes) + + @staticmethod + def _set_mtp_cuda_graph_flag(model, enabled): + """Set `use_mtp_cuda_graphs` on the model.""" + unwrapped = unwrap_model(model) + unwrapped.use_mtp_cuda_graphs = enabled + + @staticmethod + def _assert_mtp_cuda_graphs_were_replayed(model, expect_replayed): + """Assert that MTP CUDA graphs were (or were not) replayed. + + MTP runners are stored in the CudaGraphManager's lookup table + rather than the global inference record. A runner with + `fwd_graph_recorded=True` confirms the graph was captured and + replayed. + """ + unwrapped = unwrap_model(model) + manager = getattr(unwrapped, '_mtp_cudagraph_manager', None) + if manager is None: + assert not expect_replayed, "No MTP CudaGraphManager found on the model" + return + table = manager.custom_cudagraphs_lookup_table + mtp_runners = [v for k, v in table.items() if isinstance(k, tuple) and k[0] == 'mtp'] + if expect_replayed: + assert ( + len(mtp_runners) > 0 + ), "Expected MTP CUDA graphs to be replayed, but no MTP runners found" + for runner in mtp_runners: + assert runner.fwd_graph_recorded, ( + "Expected MTP CUDA graph to be recorded and replayed, " + f"but runner for {runner.base_module.__class__.__name__} " + "has fwd_graph_recorded=False" + ) + else: + recorded = [r for r in mtp_runners if r.fwd_graph_recorded] + assert len(recorded) == 0, ( + f"Expected no MTP CUDA graph replay, but {len(recorded)} " + "runners have fwd_graph_recorded=True" + ) + + # ---- Test 1: graph output matches eager (no additional padding) ------- # + + @pytest.mark.parametrize("mtp_use_repeated_layer", [False, True]) + @torch.inference_mode() + def test_cuda_graph_output_matches_eager(self, mtp_use_repeated_layer): + """CUDA graph replay produces the same output as eager execution. + + The batch sizes exactly match warmed-up graphs (from the engine's + CUDA graph warmup), so there is no additional padding. Both paths + must produce identical hidden states and logits. + """ + engine = self._build_engine(mtp_use_repeated_layer=mtp_use_repeated_layer) + model = engine.controller.inference_wrapped_model.model + unwrapped = unwrap_model(model) + batch_sizes = self._get_mtp_warmed_batch_sizes(engine) + assert len(batch_sizes) > 0, "Engine did not warm up any MTP CUDA graphs" + + mtp_depth = None if unwrapped.mtp.mtp_use_repeated_layer else 0 + + for batch_size in batch_sizes[:3]: + hidden = torch.randn( + batch_size, 1, self.HIDDEN_SIZE, device='cuda', dtype=torch.bfloat16 + ) + dist.broadcast(hidden, src=0) + token_ids = torch.randint(0, self.VOCAB_SIZE, (1, batch_size), device='cuda') + dist.broadcast(token_ids, src=0) + position_ids = torch.arange(batch_size, device='cuda', dtype=torch.int64).unsqueeze(0) + + self._set_mtp_cuda_graph_flag(model, True) + h_graph, logits_graph = unwrapped.compute_mtp_single_step( + hidden_states=hidden.clone(), + next_token_ids=token_ids.clone(), + position_ids=position_ids.clone(), + depth=mtp_depth, + ) + h_graph = h_graph.clone() + logits_graph = logits_graph.clone() + + self._set_mtp_cuda_graph_flag(model, False) + h_eager, logits_eager = unwrapped.compute_mtp_single_step( + hidden_states=hidden.clone(), + next_token_ids=token_ids.clone(), + position_ids=position_ids.clone(), + depth=mtp_depth, + ) + + torch.testing.assert_close( + h_graph, h_eager, msg=f"Hidden mismatch at batch_size={batch_size}" + ) + torch.testing.assert_close( + logits_graph, logits_eager, msg=f"Logits mismatch at batch_size={batch_size}" + ) + + self._assert_mtp_cuda_graphs_were_replayed(model, True) + + # ---- Test 2: graph matches eager with sequence parallelism ------------ # + + @pytest.mark.parametrize("mtp_use_repeated_layer", [False, True]) + @torch.inference_mode() + def test_cuda_graph_output_matches_eager_with_sp(self, mtp_use_repeated_layer): + """CUDA graph replay matches eager with sequence parallelism. + + Hidden states are in scattered SP format `[batch_size/TP, 1, H]`. + Token/position IDs remain at full `[1, batch_size]`. Both paths + must produce identical outputs. + """ + engine = self._build_engine( + sequence_parallel=True, mtp_use_repeated_layer=mtp_use_repeated_layer + ) + model = engine.controller.inference_wrapped_model.model + unwrapped = unwrap_model(model) + tp_group = parallel_state.get_tensor_model_parallel_group() + batch_sizes = self._get_mtp_warmed_batch_sizes(engine) + assert len(batch_sizes) > 0, "Engine did not warm up any MTP CUDA graphs" + + mtp_depth = None if unwrapped.mtp.mtp_use_repeated_layer else 0 + + for batch_size in batch_sizes[:3]: + hidden = torch.randn( + batch_size, 1, self.HIDDEN_SIZE, device='cuda', dtype=torch.bfloat16 + ) + dist.broadcast(hidden, src=0) + hidden_sp = scatter_to_sequence_parallel_region(hidden, group=tp_group) + + token_ids = torch.randint(0, self.VOCAB_SIZE, (1, batch_size), device='cuda') + dist.broadcast(token_ids, src=0) + position_ids = torch.arange(batch_size, device='cuda', dtype=torch.int64).unsqueeze(0) + + self._set_mtp_cuda_graph_flag(model, True) + h_graph, logits_graph = unwrapped.compute_mtp_single_step( + hidden_states=hidden_sp.clone(), + next_token_ids=token_ids.clone(), + position_ids=position_ids.clone(), + depth=mtp_depth, + ) + h_graph = h_graph.clone() + logits_graph = logits_graph.clone() + + self._set_mtp_cuda_graph_flag(model, False) + h_eager, logits_eager = unwrapped.compute_mtp_single_step( + hidden_states=hidden_sp.clone(), + next_token_ids=token_ids.clone(), + position_ids=position_ids.clone(), + depth=mtp_depth, + ) + + torch.testing.assert_close( + h_graph, h_eager, msg=f"Hidden mismatch at batch_size={batch_size}" + ) + torch.testing.assert_close( + logits_graph, logits_eager, msg=f"Logits mismatch at batch_size={batch_size}" + ) + + self._assert_mtp_cuda_graphs_were_replayed(model, True) + + # ---- Test 3: end-to-end _compute_serial_mtp_and_sample with SP ------- # + + @pytest.mark.parametrize("mtp_use_repeated_layer", [False, True]) + @torch.inference_mode() + def test_cuda_graph_sp_padding_end_to_end(self, mtp_use_repeated_layer): + """Full `_compute_serial_mtp_and_sample` with CUDA graphs and SP. + + Active request counts that are not multiples of TP are padded. + The engine's CUDA graph warmup pre-captures MTP graphs for the + padded batch sizes. Verifies that padding, SP scatter/gather, and + MTP forward all work correctly through the CUDA graph path. + """ + tp_size = self.TP_SIZE + num_spec = 2 + max_requests = 16 + engine = self._build_engine( + sequence_parallel=True, + mtp_num_layers=num_spec, + mtp_use_repeated_layer=mtp_use_repeated_layer, + num_speculative_tokens=num_spec, + max_requests=max_requests, + ) + ctrl = engine.controller + context = engine.context + model = ctrl.inference_wrapped_model.model + unwrapped = unwrap_model(model) + + mtp_sizes = self._get_mtp_warmed_batch_sizes(engine) + + # Find active_request_counts whose TP-padded values match warmed MTP sizes. + active_counts = [] + for n in mtp_sizes: + for active in range(n, 0, -1): + padded = active + (tp_size - active % tp_size) % tp_size + if padded == n and active <= max_requests: + active_counts.append(active) + break + assert len(active_counts) > 0, "No valid active request counts found" + + for active_request_count in active_counts[:4]: + padded_count = ( + active_request_count + (tp_size - active_request_count % tp_size) % tp_size + ) + + context.reset() + context.total_request_count = active_request_count + context.paused_request_count = 0 + context.request_kv_length_offsets[:active_request_count] = torch.arange( + active_request_count, dtype=torch.int32, device='cuda' + ) + context.request_query_lengths[:active_request_count] = torch.ones( + active_request_count, dtype=torch.int32, device='cuda' + ) + + ctrl.num_speculative_tokens = num_spec + ctrl.num_mtp_heads = num_spec + ctrl._init_mtp_sampling_tensors() + ctrl._mtp_token_ids_buf.zero_() + ctrl._mtp_position_ids_buf.zero_() + ctrl._sampled_tokens_cuda[:active_request_count] = torch.remainder( + torch.arange(active_request_count, device='cuda'), self.VOCAB_SIZE + ) + + tp_rank = parallel_state.get_tensor_model_parallel_rank() + + torch.manual_seed(42) + full_hidden = torch.randn( + padded_count, 1, self.HIDDEN_SIZE, device='cuda', dtype=torch.bfloat16 + ) + dist.broadcast(full_hidden, src=0) + local_hidden = full_hidden.chunk(tp_size)[tp_rank].contiguous() + unwrapped._decoder_hidden_states_cache = local_hidden + + ctrl._last_accepted_seq_indices = torch.arange(active_request_count, device='cuda') + ctrl._mtp_resolved_padded_count = padded_count + self._set_mtp_cuda_graph_flag(model, True) + + ctrl._torch_sampling_buckets = [(list(range(active_request_count)), 1.0, 1, 0.0)] + ctrl._torch_sampling_bucket_index_tensors = [ + torch.arange(active_request_count, device='cuda', dtype=torch.long) + ] + + ctrl._compute_serial_mtp_and_sample() + + for depth in range(num_spec): + sampled = ctrl._sampled_mtp_tokens_cuda[depth, :active_request_count] + assert sampled.shape == ( + active_request_count, + ), f"active={active_request_count}, depth={depth}" + assert sampled.dtype == torch.int64 + assert torch.all(sampled >= 0) and torch.all(sampled < self.VOCAB_SIZE) + + assert not hasattr(unwrapped, '_decoder_hidden_states_cache') + + self._assert_mtp_cuda_graphs_were_replayed(model, True) + + # ---- Test 4: SP padding graph vs eager produces same MTP tokens ------- # + + @pytest.mark.parametrize("mtp_use_repeated_layer", [False, True]) + @torch.inference_mode() + def test_cuda_graph_sp_padding_matches_eager(self, mtp_use_repeated_layer): + """With SP padding, CUDA graph path produces the same MTP tokens as eager. + + Uses a single engine (shared model weights) and toggles the CUDA + graph flag between runs. Both paths receive identical inputs and + must produce the same sampled MTP tokens. + """ + tp_size = self.TP_SIZE + num_spec = 2 + max_requests = 16 + engine = self._build_engine( + sequence_parallel=True, + mtp_num_layers=num_spec, + mtp_use_repeated_layer=mtp_use_repeated_layer, + num_speculative_tokens=num_spec, + max_requests=max_requests, + ) + ctrl = engine.controller + context = engine.context + model = ctrl.inference_wrapped_model.model + + mtp_sizes = self._get_mtp_warmed_batch_sizes(engine) + + # Find active counts that require TP padding (active % tp != 0). + active_counts = [] + for n in mtp_sizes: + for active in range(n, 0, -1): + padded = active + (tp_size - active % tp_size) % tp_size + if padded == n and active % tp_size != 0 and active <= max_requests: + active_counts.append(active) + break + assert len(active_counts) > 0, "No active counts with TP padding found" + + for active_request_count in active_counts[:2]: + padded_count = ( + active_request_count + (tp_size - active_request_count % tp_size) % tp_size + ) + + def _run_mtp(use_cuda_graph): + """Set up state and run MTP, returning sampled tokens.""" + unwrapped = unwrap_model(model) + context.reset() + context.total_request_count = active_request_count + context.paused_request_count = 0 + context.request_kv_length_offsets[:active_request_count] = torch.arange( + active_request_count, dtype=torch.int32, device='cuda' + ) + context.request_query_lengths[:active_request_count] = torch.ones( + active_request_count, dtype=torch.int32, device='cuda' + ) + + ctrl.num_speculative_tokens = num_spec + ctrl.num_mtp_heads = num_spec + ctrl._init_mtp_sampling_tensors() + ctrl._mtp_token_ids_buf.zero_() + ctrl._mtp_position_ids_buf.zero_() + ctrl._sampled_tokens_cuda[:active_request_count] = torch.remainder( + torch.arange(active_request_count, device='cuda'), self.VOCAB_SIZE + ) + + if use_cuda_graph: + ctrl.has_mtp_cuda_graphs = True + ctrl._mtp_resolved_padded_count = padded_count + self._set_mtp_cuda_graph_flag(model, True) + else: + ctrl.has_mtp_cuda_graphs = False + ctrl._mtp_resolved_padded_count = None + self._set_mtp_cuda_graph_flag(model, False) + + tp_rank = parallel_state.get_tensor_model_parallel_rank() + + torch.manual_seed(42) + full_hidden = torch.randn( + padded_count, 1, self.HIDDEN_SIZE, device='cuda', dtype=torch.bfloat16 + ) + dist.broadcast(full_hidden, src=0) + local_hidden = full_hidden.chunk(tp_size)[tp_rank].contiguous() + unwrapped._decoder_hidden_states_cache = local_hidden + + ctrl._last_accepted_seq_indices = torch.arange(active_request_count, device='cuda') + ctrl._torch_sampling_buckets = [(list(range(active_request_count)), 1.0, 1, 0.0)] + ctrl._torch_sampling_bucket_index_tensors = [ + torch.arange(active_request_count, device='cuda', dtype=torch.long) + ] + + ctrl._compute_serial_mtp_and_sample() + + return [ + ctrl._sampled_mtp_tokens_cuda[d, :active_request_count].clone() + for d in range(num_spec) + ] + + graph_tokens = _run_mtp(use_cuda_graph=True) + self._assert_mtp_cuda_graphs_were_replayed(model, True) + eager_tokens = _run_mtp(use_cuda_graph=False) + + for depth in range(num_spec): + assert torch.equal(graph_tokens[depth], eager_tokens[depth]), ( + f"active={active_request_count}, depth={depth}: " + f"graph tokens {graph_tokens[depth].tolist()} != " + f"eager tokens {eager_tokens[depth].tolist()}" + ) + + # ---- Test 5: multiple MTP depths with CUDA graphs --------------------- # + + @pytest.mark.parametrize("mtp_use_repeated_layer", [False, True]) + @torch.inference_mode() + def test_cuda_graph_multi_depth(self, mtp_use_repeated_layer): + """Run multiple MTP depths with CUDA graphs enabled. + + Verifies that the hidden output from one depth feeds correctly into + the next depth through the same CUDA graph, producing valid outputs + at every depth. + """ + num_depths = 2 + engine = self._build_engine( + mtp_num_layers=num_depths, mtp_use_repeated_layer=mtp_use_repeated_layer + ) + model = engine.controller.inference_wrapped_model.model + unwrapped = unwrap_model(model) + batch_sizes = self._get_mtp_warmed_batch_sizes(engine) + assert len(batch_sizes) > 0, "Engine did not warm up any MTP CUDA graphs" + + use_repeated = unwrapped.mtp.mtp_use_repeated_layer + + batch_size = batch_sizes[0] + self._set_mtp_cuda_graph_flag(model, True) + + hidden = torch.randn(batch_size, 1, self.HIDDEN_SIZE, device='cuda', dtype=torch.bfloat16) + dist.broadcast(hidden, src=0) + token_ids = torch.randint(0, self.VOCAB_SIZE, (1, batch_size), device='cuda') + dist.broadcast(token_ids, src=0) + position_ids = torch.arange(batch_size, device='cuda', dtype=torch.int64).unsqueeze(0) + + current_hidden = hidden.clone() + for depth in range(num_depths): + mtp_depth = None if use_repeated else depth + current_hidden, logits = unwrapped.compute_mtp_single_step( + hidden_states=current_hidden, + next_token_ids=token_ids.clone(), + position_ids=position_ids.clone(), + depth=mtp_depth, + ) + current_hidden = current_hidden.clone() + + assert current_hidden.shape == (batch_size, 1, self.HIDDEN_SIZE), ( + f"Depth {depth}: expected hidden shape ({batch_size}, 1, {self.HIDDEN_SIZE}), " + f"got {current_hidden.shape}" + ) + assert logits.shape == (batch_size, 1, self.VOCAB_SIZE), ( + f"Depth {depth}: expected logits shape ({batch_size}, 1, {self.VOCAB_SIZE}), " + f"got {logits.shape}" + ) + assert torch.all( + torch.isfinite(logits) + ), f"Depth {depth}: logits contain non-finite values" + + self._assert_mtp_cuda_graphs_were_replayed(model, True) + + # ---- Test 6: eager fallback when no matching graph exists ------------- # + + @pytest.mark.parametrize("mtp_use_repeated_layer", [False, True]) + @torch.inference_mode() + def test_eager_fallback_no_matching_graph(self, mtp_use_repeated_layer): + """When `use_mtp_cuda_graphs` is True but no warmed graph matches the + batch size, `compute_mtp_single_step` falls back to eager execution. + The system should produce valid outputs without errors. + """ + engine = self._build_engine(mtp_use_repeated_layer=mtp_use_repeated_layer) + model = engine.controller.inference_wrapped_model.model + unwrapped = unwrap_model(model) + warmed_sizes = set(self._get_mtp_warmed_batch_sizes(engine)) + + # Find a batch size with no matching CUDA graph. + fallback_size = None + for candidate in range(1, 32): + if candidate not in warmed_sizes: + fallback_size = candidate + break + assert fallback_size is not None, "Could not find a non-warmed batch size" + + mtp_depth = None if unwrapped.mtp.mtp_use_repeated_layer else 0 + + self._set_mtp_cuda_graph_flag(model, True) + hidden = torch.randn( + fallback_size, 1, self.HIDDEN_SIZE, device='cuda', dtype=torch.bfloat16 + ) + dist.broadcast(hidden, src=0) + token_ids = torch.randint(0, self.VOCAB_SIZE, (1, fallback_size), device='cuda') + dist.broadcast(token_ids, src=0) + position_ids = torch.arange(fallback_size, device='cuda', dtype=torch.int64).unsqueeze(0) + + h_out, logits = unwrapped.compute_mtp_single_step( + hidden_states=hidden.clone(), + next_token_ids=token_ids.clone(), + position_ids=position_ids.clone(), + depth=mtp_depth, + ) + + assert h_out.shape == (fallback_size, 1, self.HIDDEN_SIZE) + assert logits.shape == (fallback_size, 1, self.VOCAB_SIZE) + assert torch.all(torch.isfinite(logits)) + + # ---- Test 7: graph flag propagation matches main model ---------------- # + + @torch.inference_mode() + def test_mtp_graph_flag_propagation(self): + """`use_mtp_cuda_graphs` is correctly toggled via the helper.""" + model = self._build_model(mtp_num_layers=2) + unwrapped = unwrap_model(model) + + self._set_mtp_cuda_graph_flag(model, True) + assert unwrapped.use_mtp_cuda_graphs is True + + self._set_mtp_cuda_graph_flag(model, False) + assert unwrapped.use_mtp_cuda_graphs is False + + # ---- Test 8: delete_cuda_graphs resets MTP runners -------------------- # + + @torch.inference_mode() + def test_delete_cuda_graphs_resets_mtp_runners(self): + """`delete_cuda_graphs()` resets MTP CUDA graph runners. + + MTP runners are excluded from the global inference record, so they + require special handling in `delete_cuda_graphs()`. After deletion, + no MTP runners should have `fwd_graph_recorded=True` and the global + manager list should be cleared. + """ + engine = self._build_engine() + model = engine.controller.inference_wrapped_model.model + + self._assert_mtp_cuda_graphs_were_replayed(model, True) + + unwrapped = unwrap_model(model) + manager = getattr(unwrapped, '_mtp_cudagraph_manager', None) + assert manager is not None + assert len(manager.custom_cudagraphs_lookup_table) > 0 + + delete_cuda_graphs() + + assert len(manager.custom_cudagraphs_lookup_table) == 0 + assert len(_CudagraphGlobalRecord.mtp_cudagraph_managers) == 0 + + +# --------------------------------------------------------------------------- # +# TestMTPCudaGraphExpertParallel (EP = 2) +# --------------------------------------------------------------------------- # + +_EP_SIZE = 2 + +# Request state constants for parametrized tests. +NONE = "none" +DECODE = "decode" +PREFILL = "prefill" +MIXED = "mixed" + +ALL_STATES = [NONE, DECODE, PREFILL, MIXED] + +# Combinatorial sweep: C(4+2-1, 2) = 10 test cases. +_STATE_COMBOS = list(itertools.combinations_with_replacement(ALL_STATES, _EP_SIZE)) + +# Batch dimensions for each non-dummy state. +_STATE_DIMS = { + DECODE: InferenceBatchDimensions(token_count=2, prefill_req_count=0, decode_req_count=2), + PREFILL: InferenceBatchDimensions(token_count=16, prefill_req_count=2, decode_req_count=0), + MIXED: InferenceBatchDimensions(token_count=32, prefill_req_count=1, decode_req_count=2), +} + + +@pytest.mark.internal +class TestMTPCudaGraphExpertParallel: + """Tests for MTP CUDA-graphed inference with expert parallelism. + + Follows the test pattern from `test_mamba_model_expert_parallel_inference.py`. + All tests require at least `_EP_SIZE` GPUs. + """ + + HIDDEN_SIZE = 32 + VOCAB_SIZE = 100 + MAX_SEQ_LEN = 128 + NUM_LAYERS = 2 + NUM_ATTN_HEADS = 4 + NUM_MOE_EXPERTS = 2 + + @classmethod + def setup_class(cls): + if Utils.world_size < _EP_SIZE: + pytest.skip(f"EP test requires at least {_EP_SIZE} GPUs") + if Utils.world_size % _EP_SIZE != 0: + pytest.skip( + f"world_size ({Utils.world_size}) must be divisible by EP size ({_EP_SIZE})" + ) + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + expert_model_parallel_size=_EP_SIZE, + ) + + @classmethod + def teardown_class(cls): + delete_cuda_graphs() + Utils.destroy_model_parallel() + + def teardown_method(self): + delete_cuda_graphs() + + # ---- helpers ---------------------------------------------------------- # + + def _build_model(self): + """Build a GPT model with MTP + MoE + local CUDA graphs.""" + model_parallel_cuda_manual_seed(123, inference_rng_tracker=True, force_reset_rng=True) + config = TransformerConfig( + num_layers=self.NUM_LAYERS, + hidden_size=self.HIDDEN_SIZE, + num_attention_heads=self.NUM_ATTN_HEADS, + use_cpu_initialization=True, + attention_backend=AttnBackend.local, + params_dtype=torch.bfloat16, + expert_model_parallel_size=_EP_SIZE, + num_moe_experts=self.NUM_MOE_EXPERTS, + moe_token_dispatcher_type="alltoall", + add_bias_linear=False, + mtp_num_layers=2, + cuda_graph_impl="local", + moe_pad_experts_for_cuda_graph_inference=True, + ) + layer_spec = get_gpt_layer_local_spec(num_experts=self.NUM_MOE_EXPERTS) + mtp_block_spec = get_gpt_mtp_block_spec( + config=config, spec=layer_spec, use_transformer_engine=False + ) + model = GPTModel( + config=config, + transformer_layer_spec=layer_spec, + vocab_size=self.VOCAB_SIZE, + max_sequence_length=self.MAX_SEQ_LEN, + parallel_output=True, + pre_process=True, + post_process=True, + mtp_block_spec=mtp_block_spec, + ).cuda() + for param in model.parameters(): + param.data = param.data.to(config.params_dtype) + model.eval() + return model + + def _build_context( + self, + model, + *, + num_cuda_graphs=16, + use_cuda_graphs_for_non_decode_steps=True, + max_requests=None, + ): + """Build a DynamicInferenceContext for the model.""" + return DynamicInferenceContext( + model_config=model.config, + inference_config=InferenceConfig( + max_sequence_length=self.MAX_SEQ_LEN, + buffer_size_gb=0.5, + block_size_tokens=256, + materialize_only_last_token_logits=False, + num_cuda_graphs=num_cuda_graphs, + use_cuda_graphs_for_non_decode_steps=use_cuda_graphs_for_non_decode_steps, + max_requests=max_requests, + ), + ) + + # ---- Test 1: all EP ranks run MTP eager forward ----------------------- # + + @pytest.mark.parametrize("batch_size", [2, 4, 8]) + @pytest.mark.internal + @torch.inference_mode() + def test_ep_mtp_eager_forward(self, batch_size): + """All EP ranks can run MTP forward in eager mode. + + The MoE all-to-all collectives must match across EP ranks. Verifies + that all ranks complete without hanging and produce valid shapes. + """ + model = self._build_model() + unwrapped = unwrap_model(model) + + # Broadcast identical inputs so all EP ranks see the same data. + hidden = torch.randn(batch_size, 1, self.HIDDEN_SIZE, device='cuda', dtype=torch.bfloat16) + dist.broadcast(hidden, src=0) + token_ids = torch.randint(0, self.VOCAB_SIZE, (1, batch_size), device='cuda') + dist.broadcast(token_ids, src=0) + position_ids = torch.arange(batch_size, device='cuda', dtype=torch.int64).unsqueeze(0) + + h_out, logits = unwrapped.compute_mtp_single_step( + hidden_states=hidden.clone(), + next_token_ids=token_ids.clone(), + position_ids=position_ids.clone(), + depth=0, + ) + + assert h_out.shape == (batch_size, 1, self.HIDDEN_SIZE) + assert logits.shape == (batch_size, 1, self.VOCAB_SIZE) + assert torch.all(torch.isfinite(logits)) + + # ---- Test 2: dummy ranks + real ranks in eager mode ------------------- # + + @pytest.mark.internal + @torch.inference_mode() + def test_ep_mtp_eager_dummy_and_real_ranks(self): + """Even EP ranks run as dummy (with zeros), odd ranks run with real data. + + Both must issue matching MoE all-to-all collectives via the + MTP eager forward to avoid hangs. + """ + batch_size = 4 + model = self._build_model() + unwrapped = unwrap_model(model) + + ep_rank = parallel_state.get_expert_model_parallel_rank() + is_dummy = ep_rank % 2 == 0 + + if is_dummy: + hidden = torch.zeros( + batch_size, 1, self.HIDDEN_SIZE, device='cuda', dtype=torch.bfloat16 + ) + token_ids = torch.zeros(1, batch_size, device='cuda', dtype=torch.long) + else: + hidden = torch.randn( + batch_size, 1, self.HIDDEN_SIZE, device='cuda', dtype=torch.bfloat16 + ) + token_ids = torch.randint(0, self.VOCAB_SIZE, (1, batch_size), device='cuda') + position_ids = torch.arange(batch_size, device='cuda', dtype=torch.int64).unsqueeze(0) + + # All ranks must complete without hanging. + h_out, logits = unwrapped.compute_mtp_single_step( + hidden_states=hidden, next_token_ids=token_ids, position_ids=position_ids, depth=0 + ) + + assert h_out.shape == (batch_size, 1, self.HIDDEN_SIZE) + assert logits.shape == (batch_size, 1, self.VOCAB_SIZE) + + # ---- Test 3: EP state cross product with DynamicInferenceContext ------- # + + @pytest.mark.parametrize("rank_states", _STATE_COMBOS, ids=[",".join(s) for s in _STATE_COMBOS]) + @pytest.mark.internal + @torch.inference_mode() + def test_ep_state_cross_product(self, rank_states): + """Test combinatorial assignments of request states across EP ranks. + + Verifies that: + - All EP ranks agree on CUDA graph usage (on or off). + - When CUDA graphs are used, all ranks agree on the padded batch size + (which would be used as the MTP batch dimension). + """ + ep_rank = parallel_state.get_expert_model_parallel_rank() + my_state = rank_states[ep_rank] + is_dummy = my_state == NONE + + model = self._build_model() + ctx = self._build_context(model) + + # Phase 1: Set up each rank's request state. + if not is_dummy: + ctx.add_dummy_requests_for_cudagraph_capture(_STATE_DIMS[my_state]) + + # Phase 2: Initialize attention state (EP collective). + if is_dummy: + ctx.initialize_attention_state(is_expert_parallel_dummy_cuda_graph_step=True) + else: + ctx.initialize_attention_state() + + # Phase 3: Verify EP agreement on CUDA graph usage. + uses_graph = ctx.using_cuda_graph_this_step() + ep_group = parallel_state.get_expert_model_parallel_group() + uses_graph_t = torch.tensor([int(uses_graph)], device='cuda', dtype=torch.int32) + graph_min = uses_graph_t.clone() + graph_max = uses_graph_t.clone() + dist.all_reduce(graph_min, op=dist.ReduceOp.MIN, group=ep_group) + dist.all_reduce(graph_max, op=dist.ReduceOp.MAX, group=ep_group) + assert graph_min.item() == graph_max.item(), ( + f"CUDA graph usage disagrees across EP ranks: " + f"min={graph_min.item()}, max={graph_max.item()} " + f"(rank_states={rank_states})" + ) + + if not uses_graph: + return + + # Phase 4: Derive MTP padded batch size from EP-synced dimensions. + mtp_padded = ctx.padded_batch_dimensions.req_count + + # Verify MTP padded count agrees across EP ranks. + padded_t = torch.tensor([mtp_padded], dtype=torch.int32, device='cuda') + padded_max = padded_t.clone() + padded_min = padded_t.clone() + dist.all_reduce(padded_max, op=dist.ReduceOp.MAX, group=ep_group) + dist.all_reduce(padded_min, op=dist.ReduceOp.MIN, group=ep_group) + assert padded_max.item() == padded_min.item(), ( + f"MTP padded batch size mismatch across EP ranks: " + f"min={padded_min.item()}, max={padded_max.item()} " + f"(rank_states={rank_states})" + ) + + # ---- Test 4: dummy EP rank bail-out with decode-only CUDA graphs ------ # + + @pytest.mark.parametrize( + "peer_state", [PREFILL, MIXED], ids=[f"peer={s}" for s in [PREFILL, MIXED]] + ) + @pytest.mark.internal + @torch.inference_mode() + def test_ep_dummy_bailout_with_decode_only_cuda_graphs(self, peer_state): + """Verify the dummy-rank bail-out path when only decode CUDA graphs + are available. + + With `use_cuda_graphs_for_non_decode_steps=False`, only decode-only + graphs exist. When any EP rank has prefill requests, no graph matches + and all ranks fall back to eager mode. The MTP forward for the dummy + rank must use eager execution without hanging. + """ + ep_rank = parallel_state.get_expert_model_parallel_rank() + is_even = ep_rank % 2 == 0 + + model = self._build_model() + ctx = self._build_context(model, use_cuda_graphs_for_non_decode_steps=False) + + # Even ranks are dummy; odd ranks have the peer_state. + if not is_even: + ctx.add_dummy_requests_for_cudagraph_capture(_STATE_DIMS[peer_state]) + + if is_even: + ctx.initialize_attention_state(is_expert_parallel_dummy_cuda_graph_step=True) + else: + ctx.initialize_attention_state() + + # No rank should match a CUDA graph. + assert not ctx.using_cuda_graph_this_step(), ( + f"EP rank {ep_rank}: expected no CUDA graph match with " + f"decode-only graphs and peer_state={peer_state}" + ) + + # MTP eager forward should still work on all ranks. + unwrapped = unwrap_model(model) + + tp_size = parallel_state.get_tensor_model_parallel_world_size() + dummy_hidden = torch.zeros( + (tp_size, 1, self.HIDDEN_SIZE), device='cuda', dtype=torch.bfloat16 + ) + dummy_tokens = torch.zeros((1, tp_size), device='cuda', dtype=torch.long) + dummy_positions = torch.zeros((1, tp_size), device='cuda', dtype=torch.long) + + h_out, logits = unwrapped.compute_mtp_single_step( + hidden_states=dummy_hidden, + next_token_ids=dummy_tokens, + position_ids=dummy_positions, + depth=0, + ) + + assert h_out.shape == (tp_size, 1, self.HIDDEN_SIZE) + assert logits.shape == (tp_size, 1, self.VOCAB_SIZE) diff --git a/tests/unit_tests/inference/engines/test_static_engine.py b/tests/unit_tests/inference/engines/test_static_engine.py index 483a21d13bd..0067ff6e9bc 100644 --- a/tests/unit_tests/inference/engines/test_static_engine.py +++ b/tests/unit_tests/inference/engines/test_static_engine.py @@ -27,9 +27,10 @@ from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_spec from megatron.core.models.gpt.gpt_model import GPTModel from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer.cuda_graphs import delete_cuda_graphs from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import is_fa_min_version -from tests.unit_tests.test_utilities import Utils +from tests.unit_tests.test_utilities import Utils, clear_nvte_env_vars class StaticInferenceEngineTestHarness: @@ -45,11 +46,7 @@ def setup_engine( buffer_size_gb=10, inference_config_params_dtype=torch.float, ): - Utils.initialize_model_parallel( - tensor_model_parallel_size=tensor_model_parallel_size, - pipeline_model_parallel_size=pipeline_model_parallel_size, - ) - + clear_nvte_env_vars() model_parallel_cuda_manual_seed(123) self.batch_size = 4 self.hidden_size = 32 @@ -111,11 +108,23 @@ def setup_engine( buffer_size_gb=buffer_size_gb, ) + +class TestStaticInferenceEngine(StaticInferenceEngineTestHarness): + + @classmethod + def setup_class(cls): + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1 + ) + def teardown_method(self, method): - Utils.destroy_model_parallel() + delete_cuda_graphs() + @classmethod + def teardown_class(cls): + delete_cuda_graphs() + Utils.destroy_model_parallel() -class TestStaticInferenceEngine(StaticInferenceEngineTestHarness): @pytest.mark.parametrize( "batch_size,num_trials,empty_prompt", [(4, 1, False), (4, 1, True), (4, 3, False), (2, 1, False), (8, 1, False)], @@ -294,6 +303,47 @@ async def collect_stream(stream_generator, num_tokens_to_generate): f"final_streamed_token.generated_log_probs={final_streamed_token.generated_log_probs}" ) + +class TestStaticInferenceEngineParallel(StaticInferenceEngineTestHarness): + """Tests that require non-default parallel configs (varying tp/pp/ep). + + Each test initializes its own parallel state and tears it down afterward, + so these are separated from TestStaticInferenceEngine to avoid + accumulating NCCL communicator memory from repeated init/destroy cycles. + """ + + def teardown_method(self, method): + delete_cuda_graphs() + Utils.destroy_model_parallel() + + def setup_engine( + self, + engine_max_batch_size=None, + vocab_size=100, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + expert_model_parallel_size=1, + sequence_parallel=False, + legacy=False, + buffer_size_gb=10, + inference_config_params_dtype=torch.float, + ): + Utils.initialize_model_parallel( + tensor_model_parallel_size=tensor_model_parallel_size, + pipeline_model_parallel_size=pipeline_model_parallel_size, + ) + super().setup_engine( + engine_max_batch_size=engine_max_batch_size, + vocab_size=vocab_size, + tensor_model_parallel_size=tensor_model_parallel_size, + pipeline_model_parallel_size=pipeline_model_parallel_size, + expert_model_parallel_size=expert_model_parallel_size, + sequence_parallel=sequence_parallel, + legacy=legacy, + buffer_size_gb=buffer_size_gb, + inference_config_params_dtype=inference_config_params_dtype, + ) + @pytest.mark.parametrize("sequence_parallel", [False, True]) @pytest.mark.parametrize("ep_size", [1, 2]) @pytest.mark.parametrize("pp_size", [1, 2]) diff --git a/tests/unit_tests/inference/test_batch_dimension_utils.py b/tests/unit_tests/inference/test_batch_dimension_utils.py index f520c2441d7..f35be897e3a 100644 --- a/tests/unit_tests/inference/test_batch_dimension_utils.py +++ b/tests/unit_tests/inference/test_batch_dimension_utils.py @@ -122,6 +122,7 @@ class TestMatchGraphConfigWithEP: Uses the world group as the EP group (all 8 GPUs form one EP group). """ + @classmethod def setup_class(cls): Utils.initialize_model_parallel( tensor_model_parallel_size=1, @@ -129,6 +130,7 @@ def setup_class(cls): expert_model_parallel_size=Utils.world_size, ) + @classmethod def teardown_class(cls): Utils.destroy_model_parallel() @@ -352,14 +354,16 @@ def test_one_rank_oversized_forces_no_match(self, num_cuda_graphs): class TestSpeculativeDecodingBatchDimensions: """Tests for batch dimensions specifically handling speculative decoding.""" - def setup_method(self, method): + @classmethod + def setup_class(cls): Utils.initialize_model_parallel( tensor_model_parallel_size=1, pipeline_model_parallel_size=1, expert_model_parallel_size=Utils.world_size, ) - def teardown_method(self, method): + @classmethod + def teardown_class(cls): Utils.destroy_model_parallel() @staticmethod diff --git a/tests/unit_tests/inference/test_communication_utils.py b/tests/unit_tests/inference/test_communication_utils.py index 95de6c70560..e0c5a9f734d 100644 --- a/tests/unit_tests/inference/test_communication_utils.py +++ b/tests/unit_tests/inference/test_communication_utils.py @@ -1,3 +1,5 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + import pytest import torch import torch.distributed as dist @@ -22,6 +24,9 @@ def setup(self): self.size = [16, 8] self.dtype = torch.float32 + def teardown_method(self, method): + Utils.destroy_model_parallel() + @pytest.mark.skipif( not is_torch_min_version("2.4.0"), reason="torch.distributed.init_device_mesh requires torch >= 2.4.0", @@ -65,7 +70,8 @@ def test_broadcast_comparison(self, tp_size, pp_size): assert torch.allclose( tensor_received_global, tensor_received_custom ), "broadcast_from_last_pipeline_stage should be the same with or without custom pp_group" - Utils.destroy_model_parallel() + + grid.destroy() @pytest.mark.skipif( not is_torch_min_version("2.4.0"), @@ -126,4 +132,5 @@ def test_send_recv(self, tp_size, pp_size): assert torch.allclose( local_recv_buffer_global, local_recv_buffer_custom ), "Custom and global recv buffers should be the same." - Utils.destroy_model_parallel() + + grid.destroy() diff --git a/tests/unit_tests/inference/test_moe_inference.py b/tests/unit_tests/inference/test_moe_inference.py index b762b5e638c..209eab3dd83 100644 --- a/tests/unit_tests/inference/test_moe_inference.py +++ b/tests/unit_tests/inference/test_moe_inference.py @@ -10,6 +10,8 @@ - shared experts """ +import gc + import pytest import torch @@ -204,6 +206,10 @@ def teardown_class(cls): SymmetricMemoryManager.destroy() Utils.destroy_model_parallel() + def teardown_method(self, method): + gc.collect() + torch.cuda.empty_cache() + def _make_dispatcher(self, **config_overrides): from megatron.core.transformer.moe.moe_utils import get_default_pg_collection from megatron.core.transformer.moe.token_dispatcher_inference import ( diff --git a/tests/unit_tests/inference/text_generation_controllers/test_mtp_utils.py b/tests/unit_tests/inference/text_generation_controllers/test_mtp_utils.py new file mode 100644 index 00000000000..16d9d901624 --- /dev/null +++ b/tests/unit_tests/inference/text_generation_controllers/test_mtp_utils.py @@ -0,0 +1,694 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Unit tests for MTP Triton kernels. + +Each test runs both the pure-PyTorch reference (from mtp_utils_pytorch) and +the Triton kernel (from mtp_utils_triton) on the same inputs, and asserts +that the outputs match exactly. +""" + +import pytest +import torch + +from megatron.core.inference.text_generation_controllers.mtp_utils_pytorch import ( + mamba_state_selective_copy as mamba_state_selective_copy_pytorch, +) +from megatron.core.inference.text_generation_controllers.mtp_utils_pytorch import ( + prepare_next_forward_pass as prepare_next_forward_pass_pytorch, +) +from megatron.core.inference.text_generation_controllers.mtp_utils_pytorch import ( + rewind_kv_cache as rewind_kv_cache_pytorch, +) +from megatron.core.inference.text_generation_controllers.mtp_utils_pytorch import ( + verify_speculative_tokens as verify_speculative_tokens_pytorch, +) +from megatron.core.inference.text_generation_controllers.mtp_utils_triton import ( + mamba_state_selective_copy, + prepare_next_forward_pass, + rewind_kv_cache, + verify_speculative_tokens, +) + +# --------------------------------------------------------------------------- +# Test helpers +# --------------------------------------------------------------------------- + +DEVICE = "cuda" + + +def _clone_tensors(*tensors): + """Return a tuple of cloned tensors (for running reference vs kernel on the same data).""" + return tuple(t.clone() for t in tensors) + + +# --------------------------------------------------------------------------- +# Tests +# --------------------------------------------------------------------------- + + +class TestRewindKvCache: + """Tests for the rewind_kv_cache Triton kernel.""" + + @pytest.mark.parametrize("num_requests", [1, 4, 16]) + @pytest.mark.parametrize("num_speculative_tokens", [1, 2, 4]) + @pytest.mark.parametrize("block_size_tokens", [8, 16, 64]) + def test_basic(self, num_requests, num_speculative_tokens, block_size_tokens): + N = num_requests + max_blocks = 8 + + accepted_counts = torch.randint(0, num_speculative_tokens + 1, (N,), device=DEVICE) + prefill_status = torch.zeros(N, dtype=torch.int32, device=DEVICE) + + last_kv_block_offset = torch.randint(0, block_size_tokens, (N,), device=DEVICE) + kv_length_offsets = torch.randint( + block_size_tokens, block_size_tokens * 4, (N,), device=DEVICE + ) + kv_block_counts = torch.randint(2, max_blocks, (N,), device=DEVICE) + last_kv_block_id = torch.randint(0, 100, (N,), device=DEVICE) + kv_block_ids = torch.randint(0, 100, (N, max_blocks), device=DEVICE) + + ref_offset, ref_kv_len, ref_block_counts, ref_last_block, ref_block_ids = _clone_tensors( + last_kv_block_offset, kv_length_offsets, kv_block_counts, last_kv_block_id, kv_block_ids + ) + tri_offset, tri_kv_len, tri_block_counts, tri_last_block, tri_block_ids = _clone_tensors( + last_kv_block_offset, kv_length_offsets, kv_block_counts, last_kv_block_id, kv_block_ids + ) + + ref_release, ref_mask = rewind_kv_cache_pytorch( + accepted_counts.clone(), + prefill_status.clone(), + ref_offset, + ref_kv_len, + ref_block_counts, + ref_last_block, + ref_block_ids, + num_speculative_tokens, + block_size_tokens, + ) + + tri_release, tri_mask = rewind_kv_cache( + accepted_counts.clone(), + prefill_status.clone(), + tri_offset, + tri_kv_len, + tri_block_counts, + tri_last_block, + tri_block_ids, + num_speculative_tokens, + block_size_tokens, + ) + + torch.testing.assert_close(tri_offset, ref_offset) + torch.testing.assert_close(tri_kv_len, ref_kv_len) + torch.testing.assert_close(tri_block_counts, ref_block_counts) + torch.testing.assert_close(tri_last_block, ref_last_block) + torch.testing.assert_close(tri_block_ids, ref_block_ids) + torch.testing.assert_close(tri_release, ref_release) + torch.testing.assert_close(tri_mask, ref_mask) + + def test_prefill_requests_skip_rewind(self): + N = 4 + num_spec = 3 + block_size = 16 + + accepted_counts = torch.tensor([1, 0, 2, 0], device=DEVICE) + prefill_status = torch.tensor([0, 1, 0, 1], dtype=torch.int32, device=DEVICE) + last_kv_block_offset = torch.tensor([5, 10, 2, 7], device=DEVICE) + kv_length_offsets = torch.tensor([100, 200, 300, 400], device=DEVICE) + kv_block_counts = torch.tensor([3, 4, 2, 5], device=DEVICE) + last_kv_block_id = torch.tensor([10, 20, 30, 40], device=DEVICE) + kv_block_ids = torch.randint(0, 50, (N, 8), device=DEVICE) + + ref_offset, ref_kv_len, ref_block_counts, ref_last_block, ref_block_ids = _clone_tensors( + last_kv_block_offset, kv_length_offsets, kv_block_counts, last_kv_block_id, kv_block_ids + ) + tri_offset, tri_kv_len, tri_block_counts, tri_last_block, tri_block_ids = _clone_tensors( + last_kv_block_offset, kv_length_offsets, kv_block_counts, last_kv_block_id, kv_block_ids + ) + + ref_release, ref_mask = rewind_kv_cache_pytorch( + accepted_counts.clone(), + prefill_status.clone(), + ref_offset, + ref_kv_len, + ref_block_counts, + ref_last_block, + ref_block_ids, + num_spec, + block_size, + ) + tri_release, tri_mask = rewind_kv_cache( + accepted_counts.clone(), + prefill_status.clone(), + tri_offset, + tri_kv_len, + tri_block_counts, + tri_last_block, + tri_block_ids, + num_spec, + block_size, + ) + + # Prefill requests (indices 1, 3) should be unchanged. + for idx in [1, 3]: + assert ref_kv_len[idx] == kv_length_offsets[idx] + assert ref_offset[idx] == last_kv_block_offset[idx] + + torch.testing.assert_close(tri_offset, ref_offset) + torch.testing.assert_close(tri_kv_len, ref_kv_len) + torch.testing.assert_close(tri_block_counts, ref_block_counts) + torch.testing.assert_close(tri_last_block, ref_last_block) + torch.testing.assert_close(tri_block_ids, ref_block_ids) + torch.testing.assert_close(tri_mask, ref_mask) + + def test_block_boundary_crossing(self): + """When offset - rewind < 0, a block boundary is crossed.""" + N = 2 + num_spec = 3 + block_size = 16 + + accepted_counts = torch.tensor([0, 0], device=DEVICE) + prefill_status = torch.zeros(N, dtype=torch.int32, device=DEVICE) + last_kv_block_offset = torch.tensor([1, 10], device=DEVICE) + kv_length_offsets = torch.tensor([100, 200], device=DEVICE) + kv_block_counts = torch.tensor([3, 4], device=DEVICE) + last_kv_block_id = torch.tensor([50, 60], device=DEVICE) + kv_block_ids = torch.tensor( + [[10, 20, 50, -1, -1, -1, -1, -1], [15, 25, 35, 60, -1, -1, -1, -1]], device=DEVICE + ) + + ref_offset, ref_kv_len, ref_block_counts, ref_last_block, ref_block_ids = _clone_tensors( + last_kv_block_offset, kv_length_offsets, kv_block_counts, last_kv_block_id, kv_block_ids + ) + tri_offset, tri_kv_len, tri_block_counts, tri_last_block, tri_block_ids = _clone_tensors( + last_kv_block_offset, kv_length_offsets, kv_block_counts, last_kv_block_id, kv_block_ids + ) + + rewind_kv_cache_pytorch( + accepted_counts.clone(), + prefill_status.clone(), + ref_offset, + ref_kv_len, + ref_block_counts, + ref_last_block, + ref_block_ids, + num_spec, + block_size, + ) + rewind_kv_cache( + accepted_counts.clone(), + prefill_status.clone(), + tri_offset, + tri_kv_len, + tri_block_counts, + tri_last_block, + tri_block_ids, + num_spec, + block_size, + ) + + # Request 0: offset 1 - 3 = -2 → crosses boundary. + assert ref_block_counts[0] == 2 + assert tri_block_counts[0] == 2 + assert ref_last_block[0] == 20 # previous block + assert tri_last_block[0] == 20 + + # Request 1: offset 10 - 3 = 7 → no crossing. + assert ref_block_counts[1] == 4 + assert tri_block_counts[1] == 4 + + torch.testing.assert_close(tri_offset, ref_offset) + torch.testing.assert_close(tri_kv_len, ref_kv_len) + torch.testing.assert_close(tri_block_ids, ref_block_ids) + + def test_padding_programs(self): + """Padding slots (pid >= num_active_requests) must produce safe no-ops.""" + N = 8 # grid size + active = 3 + num_spec = 2 + block_size = 16 + + accepted_counts = torch.randint(0, num_spec + 1, (N,), device=DEVICE) + prefill_status = torch.zeros(N, dtype=torch.int32, device=DEVICE) + last_kv_block_offset = torch.randint(0, block_size, (N,), device=DEVICE) + kv_length_offsets = torch.randint(block_size, block_size * 4, (N,), device=DEVICE) + kv_block_counts = torch.randint(2, 6, (N,), device=DEVICE) + last_kv_block_id = torch.randint(0, 100, (N,), device=DEVICE) + kv_block_ids = torch.randint(0, 100, (N, 8), device=DEVICE) + + ref_offset, ref_kv_len, ref_block_counts, ref_last_block, ref_block_ids = _clone_tensors( + last_kv_block_offset, kv_length_offsets, kv_block_counts, last_kv_block_id, kv_block_ids + ) + tri_offset, tri_kv_len, tri_block_counts, tri_last_block, tri_block_ids = _clone_tensors( + last_kv_block_offset, kv_length_offsets, kv_block_counts, last_kv_block_id, kv_block_ids + ) + + rewind_kv_cache_pytorch( + accepted_counts.clone(), + prefill_status.clone(), + ref_offset, + ref_kv_len, + ref_block_counts, + ref_last_block, + ref_block_ids, + num_spec, + block_size, + num_active_requests=active, + ) + tri_release, tri_mask = rewind_kv_cache( + accepted_counts.clone(), + prefill_status.clone(), + tri_offset, + tri_kv_len, + tri_block_counts, + tri_last_block, + tri_block_ids, + num_spec, + block_size, + num_active_requests=active, + ) + + # Active slots should match. + torch.testing.assert_close(tri_offset[:active], ref_offset[:active]) + torch.testing.assert_close(tri_kv_len[:active], ref_kv_len[:active]) + torch.testing.assert_close(tri_block_counts[:active], ref_block_counts[:active]) + torch.testing.assert_close(tri_last_block[:active], ref_last_block[:active]) + torch.testing.assert_close(tri_block_ids[:active], ref_block_ids[:active]) + + # Padding slots: release=0, mask=False. + assert (tri_release[active:] == 0).all() + assert (~tri_mask[active:]).all() + + def test_empty(self): + N = 0 + blocks_to_release, remove_mask = rewind_kv_cache( + torch.empty(0, device=DEVICE, dtype=torch.int64), + torch.empty(0, device=DEVICE, dtype=torch.int32), + torch.empty(0, device=DEVICE, dtype=torch.int64), + torch.empty(0, device=DEVICE, dtype=torch.int64), + torch.empty(0, device=DEVICE, dtype=torch.int64), + torch.empty(0, device=DEVICE, dtype=torch.int64), + torch.empty(0, 8, device=DEVICE, dtype=torch.int64), + num_speculative_tokens=2, + block_size_tokens=16, + ) + assert blocks_to_release.shape[0] == 0 + assert remove_mask.shape[0] == 0 + + +class TestVerifySpeculativeTokens: + """Tests for the verify_speculative_tokens Triton kernel.""" + + def _make_scenario(self, num_decode, num_prefill, num_spec, *, match_pattern=None): + """Build input/output token tensors for testing. + + Args: + match_pattern: list of ints per decode request indicating how many + speculative tokens should match (0 means only base accepted). + If None, generates random matches. + """ + stride = num_spec + 1 + decode_len = num_decode * stride + total_len = decode_len + num_prefill + + input_tokens = torch.randint(1, 1000, (total_len,), device=DEVICE) + output_tokens = torch.randint(1, 1000, (total_len,), device=DEVICE) + + if match_pattern is not None: + assert len(match_pattern) == num_decode + for req_idx, num_match in enumerate(match_pattern): + base = req_idx * stride + for s in range(num_match): + output_tokens[base + s] = input_tokens[base + s + 1] + + return input_tokens, output_tokens + + @pytest.mark.parametrize( + "num_decode,num_prefill,num_spec", [(1, 0, 2), (3, 0, 2), (3, 2, 2), (0, 3, 2), (5, 3, 4)] + ) + def test_basic(self, num_decode, num_prefill, num_spec): + input_tokens, output_tokens = self._make_scenario(num_decode, num_prefill, num_spec) + + ref_last, ref_mask, ref_input = verify_speculative_tokens_pytorch( + input_tokens.clone(), output_tokens.clone(), num_decode, num_prefill, num_spec + ) + tri_last, tri_mask, tri_input = verify_speculative_tokens( + input_tokens.clone(), output_tokens.clone(), num_decode, num_prefill, num_spec + ) + + torch.testing.assert_close(tri_mask, ref_mask) + torch.testing.assert_close(tri_last, ref_last) + + def test_all_accepted(self): + """All speculative tokens match → all accepted.""" + num_decode, num_prefill, num_spec = 3, 0, 3 + input_tokens, output_tokens = self._make_scenario( + num_decode, num_prefill, num_spec, match_pattern=[3, 3, 3] + ) + + ref_last, ref_mask, _ = verify_speculative_tokens_pytorch( + input_tokens.clone(), output_tokens.clone(), num_decode, num_prefill, num_spec + ) + tri_last, tri_mask, _ = verify_speculative_tokens( + input_tokens.clone(), output_tokens.clone(), num_decode, num_prefill, num_spec + ) + + assert ref_mask.all() + torch.testing.assert_close(tri_mask, ref_mask) + torch.testing.assert_close(tri_last, ref_last) + + def test_none_accepted(self): + """No speculative tokens match → only base tokens accepted.""" + num_decode, num_prefill, num_spec = 3, 0, 3 + input_tokens, output_tokens = self._make_scenario( + num_decode, num_prefill, num_spec, match_pattern=[0, 0, 0] + ) + + ref_last, ref_mask, _ = verify_speculative_tokens_pytorch( + input_tokens.clone(), output_tokens.clone(), num_decode, num_prefill, num_spec + ) + tri_last, tri_mask, _ = verify_speculative_tokens( + input_tokens.clone(), output_tokens.clone(), num_decode, num_prefill, num_spec + ) + + stride = num_spec + 1 + for req in range(num_decode): + base = req * stride + assert ref_mask[base].item() is True + assert not ref_mask[base + 1 : base + stride].any() + + torch.testing.assert_close(tri_mask, ref_mask) + torch.testing.assert_close(tri_last, ref_last) + + def test_mixed_match_pattern(self): + """Different acceptance counts per request.""" + num_decode, num_prefill, num_spec = 3, 1, 3 + input_tokens, output_tokens = self._make_scenario( + num_decode, num_prefill, num_spec, match_pattern=[1, 3, 0] + ) + + ref_last, ref_mask, _ = verify_speculative_tokens_pytorch( + input_tokens.clone(), output_tokens.clone(), num_decode, num_prefill, num_spec + ) + tri_last, tri_mask, _ = verify_speculative_tokens( + input_tokens.clone(), output_tokens.clone(), num_decode, num_prefill, num_spec + ) + + torch.testing.assert_close(tri_mask, ref_mask) + torch.testing.assert_close(tri_last, ref_last) + + def test_2d_input(self): + """Input tokens with shape [1, total_len] should be squeezed.""" + num_decode, num_prefill, num_spec = 2, 1, 2 + input_tokens, output_tokens = self._make_scenario(num_decode, num_prefill, num_spec) + input_2d = input_tokens.unsqueeze(0) + + ref_last, ref_mask, _ = verify_speculative_tokens_pytorch( + input_2d.clone(), output_tokens.clone(), num_decode, num_prefill, num_spec + ) + tri_last, tri_mask, _ = verify_speculative_tokens( + input_2d.clone(), output_tokens.clone(), num_decode, num_prefill, num_spec + ) + + torch.testing.assert_close(tri_mask, ref_mask) + torch.testing.assert_close(tri_last, ref_last) + + +class TestPrepareNextForwardPass: + """Tests for the prepare_next_forward_pass Triton kernel.""" + + def _setup(self, num_decode, num_prefill, num_spec): + stride = num_spec + 1 + active = num_decode + num_prefill + decode_len = num_decode * stride + total_len = decode_len + num_prefill + + output_tokens = torch.randint(1, 1000, (total_len,), device=DEVICE, dtype=torch.int64) + required_logit_indices = torch.arange(total_len, device=DEVICE, dtype=torch.int64) + input_tokens = torch.randint(1, 1000, (total_len,), device=DEVICE, dtype=torch.int64) + + accepted_mask = torch.zeros(total_len, device=DEVICE, dtype=torch.bool) + last_one_indices = torch.empty(active, device=DEVICE, dtype=torch.int64) + + for req in range(num_decode): + base = req * stride + num_match = torch.randint(0, num_spec + 1, (1,)).item() + for j in range(stride): + if j <= num_match: + accepted_mask[base + j] = True + last_one_indices[req] = base + num_match + + for p in range(num_prefill): + idx = decode_len + p + accepted_mask[idx] = True + last_one_indices[num_decode + p] = idx + + return output_tokens, required_logit_indices, input_tokens, accepted_mask, last_one_indices + + @pytest.mark.parametrize( + "num_decode,num_prefill,num_spec", [(1, 0, 2), (3, 0, 2), (3, 2, 2), (0, 3, 2), (5, 3, 4)] + ) + def test_basic(self, num_decode, num_prefill, num_spec): + (output_tokens, required_logit_indices, input_tokens, accepted_mask, last_one_indices) = ( + self._setup(num_decode, num_prefill, num_spec) + ) + + active = num_decode + num_prefill + + ref_sampled = torch.zeros(active, device=DEVICE, dtype=torch.int64) + ref_last_seq = torch.zeros(active, device=DEVICE, dtype=torch.int64) + ref_accepted = torch.full((num_decode, num_spec), -1, device=DEVICE, dtype=torch.int64) + ref_counts = torch.zeros(num_decode, device=DEVICE, dtype=torch.int64) + + tri_sampled = torch.zeros(active, device=DEVICE, dtype=torch.int64) + tri_last_seq = torch.zeros(active, device=DEVICE, dtype=torch.int64) + tri_accepted = torch.full( + (max(num_decode, 1), num_spec), -1, device=DEVICE, dtype=torch.int64 + ) + tri_counts = torch.zeros(max(num_decode, 1), device=DEVICE, dtype=torch.int64) + + prepare_next_forward_pass_pytorch( + num_decode, + output_tokens, + required_logit_indices, + last_one_indices, + accepted_mask, + input_tokens, + ref_sampled, + ref_last_seq, + ref_accepted, + ref_counts, + num_spec, + ) + + prepare_next_forward_pass( + num_decode, + output_tokens, + required_logit_indices, + last_one_indices, + accepted_mask, + input_tokens, + tri_sampled, + tri_last_seq, + tri_accepted, + tri_counts, + num_spec, + ) + + torch.testing.assert_close(tri_sampled, ref_sampled) + torch.testing.assert_close(tri_last_seq, ref_last_seq) + if num_decode > 0: + torch.testing.assert_close(tri_accepted[:num_decode], ref_accepted[:num_decode]) + torch.testing.assert_close(tri_counts[:num_decode], ref_counts[:num_decode]) + + def test_empty(self): + """Zero active requests should be a no-op.""" + last_one_indices = torch.empty(0, device=DEVICE, dtype=torch.int64) + prepare_next_forward_pass( + num_decode_requests=0, + output_tokens=torch.empty(0, device=DEVICE, dtype=torch.int64), + required_logit_indices=torch.empty(0, device=DEVICE, dtype=torch.int64), + last_one_indices=last_one_indices, + accepted_tokens_mask=torch.empty(0, device=DEVICE, dtype=torch.bool), + input_tokens=torch.empty(0, device=DEVICE, dtype=torch.int64), + sampled_tokens_buf=torch.empty(0, device=DEVICE, dtype=torch.int64), + last_accepted_seq_buf=torch.empty(0, device=DEVICE, dtype=torch.int64), + accepted_tokens_per_request=torch.empty(0, 2, device=DEVICE, dtype=torch.int64), + accepted_token_counts=torch.empty(0, device=DEVICE, dtype=torch.int64), + num_speculative_tokens=2, + ) + + +class TestMambaStateSelectiveCopy: + """Tests for the mamba_state_selective_copy Triton kernel.""" + + @pytest.mark.parametrize("num_requests", [1, 4, 8]) + @pytest.mark.parametrize("num_layers", [1, 3]) + def test_basic(self, num_requests, num_layers): + N = num_requests + M = N # 1:1 request-to-slot mapping for simplicity + S = 4 # speculative tokens + 1 + state_shape = (16, 32) # arbitrary state dimensions + + intermediate = torch.randn(num_layers, M, S, *state_shape, device=DEVICE) + current_ref = torch.randn(num_layers, M, *state_shape, device=DEVICE) + current_tri = current_ref.clone() + + prefill_status = torch.zeros(N, dtype=torch.int32, device=DEVICE) + state_idx = torch.arange(N, device=DEVICE, dtype=torch.int64) + accepted_counts = torch.randint(0, S, (N,), device=DEVICE, dtype=torch.int64) + + mamba_state_selective_copy_pytorch( + intermediate, current_ref, prefill_status, state_idx, accepted_counts, num_layers + ) + mamba_state_selective_copy( + intermediate, current_tri, prefill_status, state_idx, accepted_counts, num_layers + ) + + torch.testing.assert_close(current_tri, current_ref) + + def test_prefill_skipped(self): + N = 4 + num_layers = 2 + M = N + S = 3 + state_shape = (8,) + + intermediate = torch.randn(num_layers, M, S, *state_shape, device=DEVICE) + current_ref = torch.randn(num_layers, M, *state_shape, device=DEVICE) + current_tri = current_ref.clone() + current_orig = current_ref.clone() + + prefill_status = torch.tensor([0, 1, 0, 1], dtype=torch.int32, device=DEVICE) + state_idx = torch.arange(N, device=DEVICE, dtype=torch.int64) + accepted_counts = torch.tensor([1, 0, 2, 0], device=DEVICE, dtype=torch.int64) + + mamba_state_selective_copy_pytorch( + intermediate, current_ref, prefill_status, state_idx, accepted_counts, num_layers + ) + mamba_state_selective_copy( + intermediate, current_tri, prefill_status, state_idx, accepted_counts, num_layers + ) + + # Prefill slots should be unchanged from original. + for layer in range(num_layers): + for slot in [1, 3]: + torch.testing.assert_close(current_ref[layer, slot], current_orig[layer, slot]) + torch.testing.assert_close(current_tri[layer, slot], current_orig[layer, slot]) + + torch.testing.assert_close(current_tri, current_ref) + + def test_noncontiguous_state_idx(self): + """state_idx does not have to be a simple arange.""" + N = 3 + num_layers = 2 + M = 6 # more slots than requests + S = 3 + state_shape = (8, 4) + + intermediate = torch.randn(num_layers, M, S, *state_shape, device=DEVICE) + current_ref = torch.randn(num_layers, M, *state_shape, device=DEVICE) + current_tri = current_ref.clone() + + prefill_status = torch.zeros(N, dtype=torch.int32, device=DEVICE) + state_idx = torch.tensor([1, 4, 0], device=DEVICE, dtype=torch.int64) + accepted_counts = torch.tensor([2, 0, 1], device=DEVICE, dtype=torch.int64) + + mamba_state_selective_copy_pytorch( + intermediate, current_ref, prefill_status, state_idx, accepted_counts, num_layers + ) + mamba_state_selective_copy( + intermediate, current_tri, prefill_status, state_idx, accepted_counts, num_layers + ) + + torch.testing.assert_close(current_tri, current_ref) + + def test_empty(self): + """Zero requests should be a no-op.""" + num_layers = 2 + state_shape = (8,) + intermediate = torch.randn(num_layers, 4, 3, *state_shape, device=DEVICE) + current = torch.randn(num_layers, 4, *state_shape, device=DEVICE) + current_before = current.clone() + + mamba_state_selective_copy( + intermediate, + current, + torch.empty(0, dtype=torch.int32, device=DEVICE), + torch.empty(0, dtype=torch.int64, device=DEVICE), + torch.empty(0, dtype=torch.int64, device=DEVICE), + num_layers, + ) + + torch.testing.assert_close(current, current_before) + + +class TestStressRandom: + """Randomized stress tests running all four kernels with varied inputs.""" + + @pytest.mark.parametrize("trial", range(5)) + def test_rewind_random(self, trial): + torch.manual_seed(42 + trial) + N = torch.randint(1, 32, (1,)).item() + num_spec = torch.randint(1, 6, (1,)).item() + block_size = 2 ** torch.randint(3, 7, (1,)).item() + max_blocks = torch.randint(4, 16, (1,)).item() + + accepted_counts = torch.randint(0, num_spec + 1, (N,), device=DEVICE) + prefill_status = (torch.rand(N, device=DEVICE) > 0.7).to(torch.int32) + last_kv_block_offset = torch.randint(0, block_size, (N,), device=DEVICE) + kv_length_offsets = torch.randint(block_size, block_size * 4, (N,), device=DEVICE) + kv_block_counts = torch.randint(2, max_blocks, (N,), device=DEVICE) + last_kv_block_id = torch.randint(0, 200, (N,), device=DEVICE) + kv_block_ids = torch.randint(0, 200, (N, max_blocks), device=DEVICE) + + ref_args = _clone_tensors( + last_kv_block_offset, kv_length_offsets, kv_block_counts, last_kv_block_id, kv_block_ids + ) + tri_args = _clone_tensors( + last_kv_block_offset, kv_length_offsets, kv_block_counts, last_kv_block_id, kv_block_ids + ) + + ref_release, ref_mask = rewind_kv_cache_pytorch( + accepted_counts.clone(), prefill_status.clone(), *ref_args, num_spec, block_size + ) + tri_release, tri_mask = rewind_kv_cache( + accepted_counts.clone(), prefill_status.clone(), *tri_args, num_spec, block_size + ) + + for r, t in zip(ref_args, tri_args): + torch.testing.assert_close(t, r) + torch.testing.assert_close(tri_release, ref_release) + torch.testing.assert_close(tri_mask, ref_mask) + + @pytest.mark.parametrize("trial", range(5)) + def test_verify_random(self, trial): + torch.manual_seed(42 + trial) + num_decode = torch.randint(0, 16, (1,)).item() + num_prefill = torch.randint(0, 8, (1,)).item() + if num_decode == 0 and num_prefill == 0: + num_prefill = 1 + num_spec = torch.randint(1, 6, (1,)).item() + + stride = num_spec + 1 + total_len = num_decode * stride + num_prefill + + input_tokens = torch.randint(1, 500, (total_len,), device=DEVICE) + output_tokens = torch.randint(1, 500, (total_len,), device=DEVICE) + + # Randomly make some speculative tokens match. + for req in range(num_decode): + base = req * stride + num_match = torch.randint(0, num_spec + 1, (1,)).item() + for s in range(num_match): + output_tokens[base + s] = input_tokens[base + s + 1] + + ref_last, ref_mask, _ = verify_speculative_tokens_pytorch( + input_tokens.clone(), output_tokens.clone(), num_decode, num_prefill, num_spec + ) + tri_last, tri_mask, _ = verify_speculative_tokens( + input_tokens.clone(), output_tokens.clone(), num_decode, num_prefill, num_spec + ) + + torch.testing.assert_close(tri_mask, ref_mask) + torch.testing.assert_close(tri_last, ref_last) diff --git a/tests/unit_tests/inference/text_generation_controllers/test_text_generation_controller.py b/tests/unit_tests/inference/text_generation_controllers/test_text_generation_controller.py index 9c6564f6989..dd4764ee92d 100644 --- a/tests/unit_tests/inference/text_generation_controllers/test_text_generation_controller.py +++ b/tests/unit_tests/inference/text_generation_controllers/test_text_generation_controller.py @@ -14,7 +14,7 @@ from transformer_engine.pytorch.fp8 import check_fp8_support from megatron.core import parallel_state -from megatron.core.inference.config import InferenceConfig +from megatron.core.inference.config import InferenceConfig, MambaInferenceStateConfig from megatron.core.inference.contexts import DynamicInferenceContext, StaticInferenceContext from megatron.core.inference.contexts.dynamic_context import MaxSequenceLengthOverflowError from megatron.core.inference.inference_request import ( @@ -34,6 +34,8 @@ get_gpt_mtp_block_spec, ) from megatron.core.models.gpt.gpt_model import GPTModel +from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec +from megatron.core.models.hybrid.hybrid_model import HybridModel from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.enums import AttnBackend from megatron.core.transformer.module import Float16Module @@ -43,7 +45,7 @@ from tests.unit_tests.test_utilities import Utils -class TestTextGenerationController: +class TextGenerationControllerTestBase: def setup_model( self, @@ -64,11 +66,8 @@ def setup_model( sequence_parallel: bool = False, expert_model_parallel_size: int = 1, num_moe_experts: int = None, + hybrid_layer_pattern: str = None, ): - Utils.initialize_model_parallel( - tensor_model_parallel_size=tensor_model_parallel_size, - pipeline_model_parallel_size=pipeline_model_parallel_size, - ) if use_training_random_init: # This is necessary to induce the training behavior which permutes the random seed # for every rank; otherwise, every rank will have the same seed. @@ -98,31 +97,51 @@ def setup_model( expert_model_parallel_size=expert_model_parallel_size, num_moe_experts=num_moe_experts, add_bias_linear=num_moe_experts is None, + **( + dict(is_hybrid_model=True, mamba_num_heads=2, mamba_head_dim=16, mamba_num_groups=2) + if hybrid_layer_pattern + else {} + ), ) if dtype == torch.bfloat16: transformer_config.bf16 = True - layer_spec = get_gpt_layer_local_spec() + mamba_inference_state_config = None + if hybrid_layer_pattern: + model = HybridModel( + config=transformer_config, + hybrid_stack_spec=hybrid_stack_spec, + vocab_size=self.vocab_size, + max_sequence_length=self.sequence_length, + parallel_output=True, + hybrid_layer_pattern=hybrid_layer_pattern, + pre_process=parallel_state.is_pipeline_first_stage(), + post_process=parallel_state.is_pipeline_last_stage(), + ).cuda() + mamba_inference_state_config = MambaInferenceStateConfig.from_model(model) + else: + layer_spec = get_gpt_layer_local_spec() + + mtp_block_spec = None + if mtp_num_layers > 0: + mtp_block_spec = get_gpt_mtp_block_spec( + config=transformer_config, spec=layer_spec, use_transformer_engine=False + ) - mtp_block_spec = None - if mtp_num_layers > 0: - mtp_block_spec = get_gpt_mtp_block_spec( - config=transformer_config, spec=layer_spec, use_transformer_engine=False - ) + model = GPTModel( + config=transformer_config, + transformer_layer_spec=layer_spec, + vocab_size=self.vocab_size, + max_sequence_length=self.sequence_length, + parallel_output=True, + pre_process=parallel_state.is_pipeline_first_stage(), + post_process=parallel_state.is_pipeline_last_stage(), + mtp_block_spec=mtp_block_spec, + ).cuda() - gpt_model = GPTModel( - config=transformer_config, - transformer_layer_spec=layer_spec, - vocab_size=self.vocab_size, - max_sequence_length=self.sequence_length, - parallel_output=True, - pre_process=parallel_state.is_pipeline_first_stage(), - post_process=parallel_state.is_pipeline_last_stage(), - mtp_block_spec=mtp_block_spec, - ).cuda() - gpt_model.eval() + model.eval() if dtype == torch.bfloat16: - gpt_model = Float16Module(gpt_model.config, gpt_model) + model = Float16Module(model.config, model) if static: inference_context = StaticInferenceContext( @@ -142,10 +161,11 @@ def setup_model( block_size_tokens=block_size_tokens, enable_prefix_caching=enable_prefix_caching, max_requests=max_requests, + mamba_inference_state_config=mamba_inference_state_config, ), ) - inference_wrapped_model = GPTInferenceWrapper(gpt_model, inference_context) + inference_wrapped_model = GPTInferenceWrapper(model, inference_context) inference_wrapped_model.model_is_pipeline_parallel = not ( parallel_state.is_pipeline_first_stage() and parallel_state.is_pipeline_last_stage() @@ -157,6 +177,15 @@ def setup_model( inference_wrapped_model=inference_wrapped_model, tokenizer=self.mock_tokenizer ) + +class TestTextGenerationController(TextGenerationControllerTestBase): + + @classmethod + def setup_class(cls): + Utils.initialize_model_parallel( + tensor_model_parallel_size=2, pipeline_model_parallel_size=1 + ) + @classmethod def teardown_class(cls): Utils.destroy_model_parallel() @@ -931,164 +960,12 @@ def test_dynamic_top_n_logprobs_calculation( top_n_indices.shape[0] == top_n ), f"Request {req_idx}, token {token_idx}: expected {top_n} indices" - @pytest.mark.parametrize("static", [True, False]) - @pytest.mark.parametrize("tp_size", [1, 2]) - @pytest.mark.parametrize("pp_size", [1, 2]) - def test_sampled_tokens_match_with_parallelism(self, static, tp_size, pp_size): - """ - Verify that sampled tokens match across all parallel ranks. - """ - if tp_size == 1 and pp_size == 1: - pytest.skip(reason="Test requires model parallel size > 1.") - - if not static and not is_fa_min_version("2.7.3"): - pytest.skip(reason="Need latest flash attn for dynamic batching") - - # Ensure that we are using the training setup for random seed initialization - # so that every rank has a different seed - self.setup_model( - dtype=torch.bfloat16, - tensor_model_parallel_size=tp_size, - pipeline_model_parallel_size=pp_size, - static=static, - use_training_random_init=True, - ) - - self.mock_tokenizer.vocab_size = self.vocab_size - self.mock_tokenizer.eod = self.vocab_size - 1 - self.mock_tokenizer.detokenize.side_effect = lambda x, skip_special_tokens=False: ' '.join( - [ - ''.join(random.choices(string.ascii_letters, k=random.randint(4, 10))) - for _ in range(len(x)) - ] - ) - self.mock_tokenizer.offsets.side_effect = lambda _, s: [ - i for i, c in enumerate(s) if c == ' ' - ] + [len(s)] - - # Prepare requests. - active_requests: Dict[str, InferenceRequest] = OrderedDict() - for i in range(self.batch_size): - prompt = "sample" * (i + 1) - prompt_tokens = torch.randint( - low=0, high=self.vocab_size - 1, size=(len(prompt),) - ).tolist() - request_id = str(i) - inference_request = InferenceRequest( - request_id=request_id, - prompt=prompt, - sampling_params=SamplingParams( - top_k=10, num_tokens_to_generate=25, return_log_probs=True - ), - arrival_time=time.time(), - prompt_tokens=prompt_tokens, - status=Status.ACTIVE_BUT_NOT_GENERATING_TOKENS, - ) - active_requests[request_id] = inference_request - - # Generate tokens. - if static: - requests = self.text_generation_controller.generate_all_output_tokens_static_batch( - active_requests - ) - all_generated_tokens = [req.generated_tokens.tolist() for req in requests.values()] - else: - all_generated_tokens = [[] for _ in range(len(active_requests))] - context = self.text_generation_controller.inference_wrapped_model.inference_context - for request_id, request in active_requests.items(): - context.add_request( - DynamicInferenceRequest( - request_id=int(request_id), - prompt_tokens=torch.tensor( - request.prompt_tokens, - dtype=torch.long, - device=torch.cuda.current_device(), - ), - sampling_params=SamplingParams( - top_k=10, return_log_probs=True, num_tokens_to_generate=25 - ), - ) - ) - expected_active_requests = set(int(x) for x in active_requests.keys()) - while context.has_unfinished_requests(): - result = self.text_generation_controller.generate_output_tokens_dynamic_batch() - new_tokens = result["sample"] - active_ids = result["active_request_ids"].tolist() - finished_ids = result["finished_request_ids"].tolist() - assert len(new_tokens) == len(expected_active_requests) - assert set(active_ids) == expected_active_requests - expected_active_requests -= set(finished_ids) - for i, token in enumerate(new_tokens.tolist()): - all_generated_tokens[i].append(token) - - # Wait for all communication to complete before proceeding. - torch.distributed.barrier() - - # Collect all the generated tokens for each request from each rank in the - # model parallel group. - mp_group = parallel_state.get_model_parallel_group() - mp_ranks = torch.distributed.get_process_group_ranks(mp_group) - local_rank = torch.distributed.get_rank() - tokens_per_rank = {} - tokens_per_rank[local_rank] = all_generated_tokens - - for i in mp_ranks: - # Start by communicating the batch size so each rank knows how many requests to expect. - if i == local_rank: - batch_size = torch.tensor( - len(tokens_per_rank[local_rank]), - dtype=torch.long, - device=torch.cuda.current_device(), - ) - else: - tokens_per_rank[i] = [] - batch_size = torch.empty(1, dtype=torch.long, device=torch.cuda.current_device()) - torch.distributed.broadcast(batch_size, group=mp_group, src=i) - - for j in range(batch_size.item()): - # For each request, communicate the sequence length followed by the actual tokens. - if i == local_rank: - sequence_length = torch.tensor( - len(tokens_per_rank[local_rank][j]), - dtype=torch.int32, - device=torch.cuda.current_device(), - ) - else: - sequence_length = torch.empty( - 1, dtype=torch.int32, device=torch.cuda.current_device() - ) - torch.distributed.broadcast(sequence_length, group=mp_group, src=i) - - if i == local_rank: - generated_tokens = torch.tensor( - tokens_per_rank[local_rank][j], - dtype=torch.long, - device=torch.cuda.current_device(), - ) - else: - generated_tokens = torch.empty( - sequence_length.item(), dtype=torch.long, device=torch.cuda.current_device() - ) - torch.distributed.broadcast(generated_tokens, group=mp_group, src=i) - - if i != local_rank: - tokens_per_rank[i].append(generated_tokens.tolist()) - - # Ensure that every rank in the model parallel group produced the same tokens. - for i in mp_ranks: - if i == local_rank: - continue - for j, (expected, actual) in enumerate( - zip(tokens_per_rank[local_rank], tokens_per_rank[i]) - ): - assert ( - expected == actual - ), f"Rank {i} tokens differ from rank {local_rank} tokens for request {j}" - @pytest.mark.internal def test_speculative_verify_tokens(self): """Test consecutive token acceptance logic for speculative decoding.""" - self.setup_model(torch.float32, static=False, num_speculative_tokens=2, max_requests=2) + self.setup_model( + torch.float32, static=False, num_speculative_tokens=2, max_requests=2, mtp_num_layers=2 + ) # Enable speculative decoding self.text_generation_controller.num_speculative_tokens = 2 @@ -1103,7 +980,7 @@ def test_speculative_verify_tokens(self): ) # 1 sampled + 2 spec # Init accepted tokens tensors - self.text_generation_controller._init_mtp_sampling_tensor() + self.text_generation_controller._init_mtp_sampling_tensors() # Mock inputs: [Req 1 sampled, Req 1 spec1, Req 1 spec2, Req 2 sampled, Req 2 spec1, Req 2 spec2] # Target tokens (what the model was fed): [T0, T1, T2, T3, T4, T5] @@ -1124,6 +1001,9 @@ def mock_sampling_func(logits, *args, **kwargs): # Override sampling to return our predictable mock outputs self.text_generation_controller._torch_sampling_buckets = [([0, 1], 1.0, 1, 0.0)] + self.text_generation_controller._torch_sampling_bucket_index_tensors = [ + torch.tensor([0, 1], device='cuda', dtype=torch.long) + ] self.text_generation_controller._torch_sampling_func = mock.MagicMock( side_effect=mock_sampling_func ) @@ -1156,6 +1036,7 @@ def test_rewind_kv_cache(self, is_hybrid_model): num_speculative_tokens=3, block_size_tokens=4, max_requests=16, + hybrid_layer_pattern="***M" if is_hybrid_model else None, ) self.text_generation_controller.num_speculative_tokens = 3 ctx = self.text_generation_controller.inference_wrapped_model.inference_context @@ -1177,21 +1058,22 @@ def test_rewind_kv_cache(self, is_hybrid_model): ) if is_hybrid_model: - ctx.is_hybrid_model = True - ctx.mamba_metadata = mock.MagicMock() - ctx.mamba_metadata.request_to_mamba_state_idx = torch.tensor([0, 1], device='cuda') - ctx.mamba_ssm_states = torch.zeros((1, 2, 16), device='cuda') - ctx.mamba_intermediate_ssm_states = torch.ones((1, 2, 4, 16), device='cuda') * 99 - ctx.mamba_conv_states = torch.zeros((1, 2, 8), device='cuda') - ctx.mamba_intermediate_conv_states = torch.ones((1, 2, 4, 8), device='cuda') * 77 + ctx.mamba_metadata.request_to_mamba_state_idx[:2] = torch.tensor( + [0, 1], dtype=torch.int32, device='cuda' + ) + ctx.mamba_ssm_states.zero_() + ctx.mamba_intermediate_ssm_states.fill_(99) + ctx.mamba_conv_states.zero_() + ctx.mamba_intermediate_conv_states.fill_(77) # Mock accepted token counts: Req 0 accepts 1 (rejects 2), Req 1 accepts 0 (rejects 3) - self.text_generation_controller._init_mtp_sampling_tensor() + self.text_generation_controller._init_mtp_sampling_tensors() self.text_generation_controller._accepted_token_counts_per_request = torch.tensor( [1, 0], device='cuda' ) - self.text_generation_controller._rewind_kv_cache() + blocks_to_release, remove_mask = self.text_generation_controller._rewind_kv_cache() + ctx.kv_block_allocator.release_memory_blocks(blocks_to_release[remove_mask]) # Assert offsets updated assert torch.equal( @@ -1221,6 +1103,108 @@ def test_rewind_kv_cache(self, is_hybrid_model): assert torch.all(ctx.mamba_conv_states[:, 0] == 77) # Req 0 accepted 1, loaded index 1 assert torch.all(ctx.mamba_conv_states[:, 1] == 77) # Req 1 accepted 0, loaded index 0 + @pytest.mark.internal + def test_rewind_kv_cache_stale_padding_is_safe(self): + """Padding slots with stale data must not corrupt active requests or + release junk blocks when the rewind kernel grid is padded beyond the + active request count. + + Without the num_active_requests guard in the kernel, padding slots + whose stale request_last_kv_block_offset < num_speculative_tokens + would produce remove_mask=True, causing the block allocator to free + block IDs that belong to other active requests. + """ + from megatron.core.inference.text_generation_controllers.mtp_utils_triton import ( + rewind_kv_cache, + ) + + num_spec = 3 + block_size = 4 + active = 2 + padded = 4 + max_blocks = 10 + dev = 'cuda' + + # --- Active requests (slots 0-1): identical to test_rewind_kv_cache --- + # Req 0: accepted 1, last_offset 2 → rewind 2 → offset 0, no release + # Req 1: accepted 0, last_offset 1 → rewind 3 → crosses block, release block 60 + accepted = torch.zeros(padded, device=dev, dtype=torch.int64) + accepted[0] = 1 + accepted[1] = 0 + + prefill = torch.zeros(padded, device=dev, dtype=torch.int64) + + last_offset = torch.zeros(padded, device=dev, dtype=torch.int64) + last_offset[0] = 2 + last_offset[1] = 1 + + kv_length = torch.zeros(padded, device=dev, dtype=torch.int64) + kv_length[0] = 10 + kv_length[1] = 15 + + block_counts = torch.zeros(padded, device=dev, dtype=torch.int64) + block_counts[0] = 3 + block_counts[1] = 4 + + last_block_id = torch.zeros(padded, device=dev, dtype=torch.int64) + last_block_id[0] = 50 + last_block_id[1] = 60 + + block_ids = torch.full((padded, max_blocks), -1, device=dev, dtype=torch.int64) + block_ids[0, :3] = torch.tensor([48, 49, 50]) + block_ids[1, :4] = torch.tensor([57, 58, 59, 60]) + + # --- Padding slots (2-3): stale data from completed requests --- + # Crucially, last_offset values < num_spec would trigger remove=True + # without the kernel guard, releasing stale block IDs. + last_offset[2] = 1 + last_offset[3] = 2 + kv_length[2] = 9999 + kv_length[3] = 9999 + block_counts[2] = 5 + block_counts[3] = 7 + last_block_id[2] = 777 + last_block_id[3] = 888 + block_ids[2, :5] = torch.arange(100, 105, device=dev) + block_ids[3, :5] = torch.arange(200, 205, device=dev) + + blocks_to_release, remove_mask = rewind_kv_cache( + accepted_counts=accepted, + prefill_status=prefill, + last_kv_block_offset=last_offset, + kv_length_offsets=kv_length, + kv_block_counts=block_counts, + last_kv_block_id=last_block_id, + kv_block_ids=block_ids, + num_speculative_tokens=num_spec, + block_size_tokens=block_size, + num_active_requests=active, + ) + + # --- Active request 0: rewind 2, no block release --- + assert remove_mask[0].item() is False + assert last_offset[0].item() == 0 + assert kv_length[0].item() == 8 + assert block_counts[0].item() == 3 + assert last_block_id[0].item() == 50 + + # --- Active request 1: rewind 3, crosses block boundary --- + assert remove_mask[1].item() is True + assert last_offset[1].item() == 2 # (1 - 3) % 4 = 2 + assert kv_length[1].item() == 12 + assert block_counts[1].item() == 3 + assert last_block_id[1].item() == 59 + assert blocks_to_release[1].item() == 60 + + # --- Padding slots 2-3: must be no-ops, no blocks released --- + assert remove_mask[2].item() is False + assert remove_mask[3].item() is False + # Stale state must be untouched (kernel skipped these programs). + assert kv_length[2].item() == 9999 + assert kv_length[3].item() == 9999 + assert block_counts[2].item() == 5 + assert block_counts[3].item() == 7 + @pytest.mark.internal def test_speculative_multinomial_sampling(self): """Test that speculative decoding can successfully use non-greedy sampling @@ -1251,6 +1235,9 @@ def test_speculative_multinomial_sampling(self): # Set up a bucket that forces multinomial sampling (top_p = 0.9, top_k = 0) # _torch_sampling_buckets format: (indices, temp, top_k, top_p) self.text_generation_controller._torch_sampling_buckets = [([0, 1], 1.0, 0, 0.9)] + self.text_generation_controller._torch_sampling_bucket_index_tensors = [ + torch.tensor([0, 1], device='cuda', dtype=torch.long) + ] # Since we are actually testing the internal math of `_torch_sampling_func` handling the shapes, # we DO NOT mock `_torch_sampling_func` here. We want it to run natively to prove it doesn't crash. @@ -1310,12 +1297,13 @@ def test_rewind_kv_cache_with_prefix_caching_ref_counts(self): initial_avail = ctx.kv_block_allocator.total_avail # Req 0 accepts 1 (rewinds 1), Req 1 accepts 0 (rewinds 2, crosses boundary). - self.text_generation_controller._init_mtp_sampling_tensor() + self.text_generation_controller._init_mtp_sampling_tensors() self.text_generation_controller._accepted_token_counts_per_request = torch.tensor( [1, 0], device='cuda' ) - self.text_generation_controller._rewind_kv_cache() + blocks_to_release, remove_mask = self.text_generation_controller._rewind_kv_cache() + ctx.kv_block_allocator.release_memory_blocks(blocks_to_release[remove_mask]) # Req 1 should have released block 20 (ref count decremented). assert ctx.kv_block_allocator.block_ref_counts[20].item() == 1 @@ -1350,12 +1338,13 @@ def test_rewind_kv_cache_does_not_release_shared_prefix_blocks(self): # Blocks 10, 20 are shared prefix blocks. Block 30, 40 are exclusive. ctx.kv_block_allocator.total_avail = 50 - self.text_generation_controller._init_mtp_sampling_tensor() + self.text_generation_controller._init_mtp_sampling_tensors() self.text_generation_controller._accepted_token_counts_per_request = torch.tensor( [0], device='cuda' ) - self.text_generation_controller._rewind_kv_cache() + blocks_to_release, remove_mask = self.text_generation_controller._rewind_kv_cache() + ctx.kv_block_allocator.release_memory_blocks(blocks_to_release[remove_mask]) # Only block 40 should be released, not blocks 10, 20, or 30. assert ctx.request_kv_block_counts[0].item() == 3 @@ -1372,7 +1361,9 @@ def test_rewind_kv_cache_does_not_release_shared_prefix_blocks(self): def test_speculative_mtp_position_ids_with_prefill(self): """Test that _compute_serial_mtp_and_sample uses the correct position IDs for a mixed batch of prefill and decode requests.""" - self.setup_model(torch.float32, static=False, num_speculative_tokens=2, max_requests=2) + self.setup_model( + torch.float32, static=False, num_speculative_tokens=2, max_requests=2, mtp_num_layers=2 + ) self.text_generation_controller.num_speculative_tokens = 2 self.text_generation_controller.num_mtp_heads = 2 @@ -1388,7 +1379,7 @@ def test_speculative_mtp_position_ids_with_prefill(self): ctx.request_kv_length_offsets[:2] = torch.tensor([10, 0], dtype=torch.int32, device='cuda') ctx.request_query_lengths[:2] = torch.tensor([3, 15], dtype=torch.int32, device='cuda') - self.text_generation_controller._init_mtp_sampling_tensor() + self.text_generation_controller._init_mtp_sampling_tensors() # Mock base token sampling (the first tokens fed into MTP) self.text_generation_controller._sampled_tokens_cuda[:2] = torch.tensor( [100, 200], device='cuda' @@ -1403,7 +1394,7 @@ def test_speculative_mtp_position_ids_with_prefill(self): captured_position_ids = [] - def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth): + def mock_compute_mtp_single_step(hidden_states, next_token_ids, position_ids, depth=None): captured_position_ids.append(position_ids.clone()) return hidden_states, torch.randn(2, 1, self.vocab_size, device='cuda') @@ -1463,7 +1454,7 @@ def test_mtp_sp_padding_real_ranks(self, active_request_count): active_request_count, dtype=torch.int32, device='cuda' ) - ctrl._init_mtp_sampling_tensor() + ctrl._init_mtp_sampling_tensors() ctrl._sampled_tokens_cuda[:active_request_count] = torch.remainder( torch.arange(active_request_count, device='cuda'), self.vocab_size ) @@ -1488,6 +1479,9 @@ def test_mtp_sp_padding_real_ranks(self, active_request_count): # Greedy sampling: top_k=1 selects the argmax token deterministically. ctrl._torch_sampling_buckets = [(list(range(active_request_count)), 1.0, 1, 0.0)] + ctrl._torch_sampling_bucket_index_tensors = [ + torch.arange(active_request_count, device='cuda', dtype=torch.long) + ] # Run the MTP forward pass ctrl._compute_serial_mtp_and_sample() @@ -1539,10 +1533,7 @@ def test_mtp_sp_padding_dummy_ranks(self): dummy_positions = torch.zeros((1, tp_size), device='cuda', dtype=torch.long) hidden_out, logits_out = unwrapped_model.compute_mtp_single_step( - hidden_states=dummy_hidden, - next_token_ids=dummy_tokens, - position_ids=dummy_positions, - depth=0, + hidden_states=dummy_hidden, next_token_ids=dummy_tokens, position_ids=dummy_positions ) # Hidden output is in SP format: [padded_count/tp_size, 1, H] = [1, 1, H]. @@ -1585,7 +1576,6 @@ def test_mtp_sp_dummy_hidden_uses_full_seq_len(self): hidden_states=current_hidden, next_token_ids=dummy_tokens, position_ids=dummy_positions, - depth=depth, ) # Hidden stays in SP format across all depths. @@ -1598,3 +1588,209 @@ def test_mtp_sp_dummy_hidden_uses_full_seq_len(self): f"Depth {depth}: expected logits shape ({tp_size}, 1, {self.vocab_size}), " f"got {logits.shape}" ) + + +class TestTextGenerationControllerParallel(TextGenerationControllerTestBase): + """Tests that require non-default parallel configs (varying tp/pp). + + Each test initializes its own parallel state and tears it down afterward, + so these are separated from TestTextGenerationController to avoid + accumulating NCCL communicator memory from repeated init/destroy cycles. + """ + + def teardown_method(self, method): + Utils.destroy_model_parallel() + + def setup_model( + self, + dtype, + symmetric_ar_type=None, + fp8: bool = False, + tensor_model_parallel_size: int = 2, + pipeline_model_parallel_size: int = 1, + batch_size: int = 4, + static: bool = True, + use_training_random_init: bool = False, + materialize_only_last_token_logits: bool = False, + num_speculative_tokens: int = 0, + block_size_tokens: int = 256, + enable_prefix_caching: bool = False, + max_requests: int = None, + mtp_num_layers: int = 0, + sequence_parallel: bool = False, + expert_model_parallel_size: int = 1, + num_moe_experts: int = None, + hybrid_layer_pattern: str = None, + ): + Utils.initialize_model_parallel( + tensor_model_parallel_size=tensor_model_parallel_size, + pipeline_model_parallel_size=pipeline_model_parallel_size, + ) + super().setup_model( + dtype, + symmetric_ar_type=symmetric_ar_type, + fp8=fp8, + tensor_model_parallel_size=tensor_model_parallel_size, + pipeline_model_parallel_size=pipeline_model_parallel_size, + batch_size=batch_size, + static=static, + use_training_random_init=use_training_random_init, + materialize_only_last_token_logits=materialize_only_last_token_logits, + num_speculative_tokens=num_speculative_tokens, + block_size_tokens=block_size_tokens, + enable_prefix_caching=enable_prefix_caching, + max_requests=max_requests, + mtp_num_layers=mtp_num_layers, + sequence_parallel=sequence_parallel, + expert_model_parallel_size=expert_model_parallel_size, + num_moe_experts=num_moe_experts, + hybrid_layer_pattern=hybrid_layer_pattern, + ) + + @pytest.mark.parametrize("static", [True, False]) + @pytest.mark.parametrize("tp_size", [1, 2]) + @pytest.mark.parametrize("pp_size", [1, 2]) + def test_sampled_tokens_match_with_parallelism(self, static, tp_size, pp_size): + """Verify that sampled tokens match across all parallel ranks.""" + if tp_size == 1 and pp_size == 1: + pytest.skip(reason="Test requires model parallel size > 1.") + + if not static and not is_fa_min_version("2.7.3"): + pytest.skip(reason="Need latest flash attn for dynamic batching") + + self.setup_model( + dtype=torch.bfloat16, + tensor_model_parallel_size=tp_size, + pipeline_model_parallel_size=pp_size, + static=static, + use_training_random_init=True, + ) + + self.mock_tokenizer.vocab_size = self.vocab_size + self.mock_tokenizer.eod = self.vocab_size - 1 + self.mock_tokenizer.detokenize.side_effect = lambda x, skip_special_tokens=False: ' '.join( + [ + ''.join(random.choices(string.ascii_letters, k=random.randint(4, 10))) + for _ in range(len(x)) + ] + ) + self.mock_tokenizer.offsets.side_effect = lambda _, s: [ + i for i, c in enumerate(s) if c == ' ' + ] + [len(s)] + + # Prepare requests. + active_requests: Dict[str, InferenceRequest] = OrderedDict() + for i in range(self.batch_size): + prompt = "sample" * (i + 1) + prompt_tokens = torch.randint( + low=0, high=self.vocab_size - 1, size=(len(prompt),) + ).tolist() + request_id = str(i) + inference_request = InferenceRequest( + request_id=request_id, + prompt=prompt, + sampling_params=SamplingParams( + top_k=10, num_tokens_to_generate=25, return_log_probs=True + ), + arrival_time=time.time(), + prompt_tokens=prompt_tokens, + status=Status.ACTIVE_BUT_NOT_GENERATING_TOKENS, + ) + active_requests[request_id] = inference_request + + # Generate tokens. + if static: + requests = self.text_generation_controller.generate_all_output_tokens_static_batch( + active_requests + ) + all_generated_tokens = [req.generated_tokens.tolist() for req in requests.values()] + else: + all_generated_tokens = [[] for _ in range(len(active_requests))] + context = self.text_generation_controller.inference_wrapped_model.inference_context + for request_id, request in active_requests.items(): + context.add_request( + DynamicInferenceRequest( + request_id=int(request_id), + prompt_tokens=torch.tensor( + request.prompt_tokens, + dtype=torch.long, + device=torch.cuda.current_device(), + ), + sampling_params=SamplingParams( + top_k=10, return_log_probs=True, num_tokens_to_generate=25 + ), + ) + ) + expected_active_requests = set(int(x) for x in active_requests.keys()) + while context.has_unfinished_requests(): + result = self.text_generation_controller.generate_output_tokens_dynamic_batch() + new_tokens = result["sample"] + active_ids = result["active_request_ids"].tolist() + finished_ids = result["finished_request_ids"].tolist() + assert len(new_tokens) == len(expected_active_requests) + assert set(active_ids) == expected_active_requests + expected_active_requests -= set(finished_ids) + for i, token in enumerate(new_tokens.tolist()): + all_generated_tokens[i].append(token) + + # Wait for all communication to complete before proceeding. + torch.distributed.barrier() + + # Collect all the generated tokens for each request from each rank in the + # model parallel group. + mp_group = parallel_state.get_model_parallel_group() + mp_ranks = torch.distributed.get_process_group_ranks(mp_group) + local_rank = torch.distributed.get_rank() + tokens_per_rank = {} + tokens_per_rank[local_rank] = all_generated_tokens + + for i in mp_ranks: + if i == local_rank: + batch_size = torch.tensor( + len(tokens_per_rank[local_rank]), + dtype=torch.long, + device=torch.cuda.current_device(), + ) + else: + tokens_per_rank[i] = [] + batch_size = torch.empty(1, dtype=torch.long, device=torch.cuda.current_device()) + torch.distributed.broadcast(batch_size, group=mp_group, src=i) + + for j in range(batch_size.item()): + if i == local_rank: + sequence_length = torch.tensor( + len(tokens_per_rank[local_rank][j]), + dtype=torch.int32, + device=torch.cuda.current_device(), + ) + else: + sequence_length = torch.empty( + 1, dtype=torch.int32, device=torch.cuda.current_device() + ) + torch.distributed.broadcast(sequence_length, group=mp_group, src=i) + + if i == local_rank: + generated_tokens = torch.tensor( + tokens_per_rank[local_rank][j], + dtype=torch.long, + device=torch.cuda.current_device(), + ) + else: + generated_tokens = torch.empty( + sequence_length.item(), dtype=torch.long, device=torch.cuda.current_device() + ) + torch.distributed.broadcast(generated_tokens, group=mp_group, src=i) + + if i != local_rank: + tokens_per_rank[i].append(generated_tokens.tolist()) + + # Ensure that every rank in the model parallel group produced the same tokens. + for i in mp_ranks: + if i == local_rank: + continue + for j, (expected, actual) in enumerate( + zip(tokens_per_rank[local_rank], tokens_per_rank[i]) + ): + assert ( + expected == actual + ), f"Rank {i} tokens differ from rank {local_rank} tokens for request {j}" diff --git a/tests/unit_tests/test_utilities.py b/tests/unit_tests/test_utilities.py index f8fad3325f5..0ddfef4dc67 100644 --- a/tests/unit_tests/test_utilities.py +++ b/tests/unit_tests/test_utilities.py @@ -26,6 +26,13 @@ def __init__( self.layers[-1].weight.shared_embedding = True +def clear_nvte_env_vars(): + """Clear NVTE env vars set by conftest set_env fixture.""" + os.environ.pop('NVTE_FLASH_ATTN', None) + os.environ.pop('NVTE_FUSED_ATTN', None) + os.environ.pop('NVTE_UNFUSED_ATTN', None) + + class Utils: world_size = int(os.environ.get('WORLD_SIZE', '1'))