Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
27 commits
Select commit Hold shift + click to select a range
8f35b31
Make MTP / prefix cache states persist for engine lifetime
santhnm2 Apr 1, 2026
e16d0e2
Merge with main
santhnm2 Apr 9, 2026
917dc1b
Track per-token acceptance rates
santhnm2 Apr 9, 2026
123cac4
Merge with main
santhnm2 May 20, 2026
78f0ce9
Fix looping
santhnm2 May 20, 2026
529f04f
Merge branch 'main' into mtp_prefix_caching_stats
santhnm2 May 20, 2026
add45a4
Remove unnecessary line
santhnm2 May 20, 2026
f68ccc3
Replace num_mtp_heads with num_mtp_depths computed from model config
santhnm2 May 21, 2026
d9d00e1
Address reviewer comment
santhnm2 May 21, 2026
1f5f73d
Linting
santhnm2 May 21, 2026
e9d4147
Merge branch 'main' into mtp_prefix_caching_stats
santhnm2 May 21, 2026
a5cb845
Return accepted_tokens=None for prefill-only batches in speculative d…
santhnm2 May 22, 2026
6439fa0
Fix extra kwargs
santhnm2 May 22, 2026
f08ed87
Linting
santhnm2 May 22, 2026
52e65c5
Vectorize
santhnm2 May 22, 2026
777572b
Linting
santhnm2 May 22, 2026
6a87249
Merge branch 'main' into mtp_prefix_caching_stats
santhnm2 May 22, 2026
edb3957
Support full model cuda graphs
santhnm2 May 22, 2026
eb4c646
Revert "Support full model cuda graphs"
santhnm2 May 22, 2026
ea634f9
Merge branch 'main' into mtp_prefix_caching_stats
santhnm2 Jun 1, 2026
491f263
Merge branch 'main' into mtp_prefix_caching_stats
santhnm2 Jun 3, 2026
3fcd581
Merge branch 'main' into mtp_prefix_caching_stats
santhnm2 Jun 4, 2026
4f9ee47
Comment fixes
santhnm2 Jun 4, 2026
6431116
Fix unit tests
santhnm2 Jun 4, 2026
40409f0
Merge branch 'main' into mtp_prefix_caching_stats
santhnm2 Jun 4, 2026
d1dff10
Fix tests
santhnm2 Jun 4, 2026
41b6d99
Merge branch 'main' into mtp_prefix_caching_stats
santhnm2 Jun 4, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
9 changes: 5 additions & 4 deletions megatron/core/inference/contexts/dynamic_context.py
Original file line number Diff line number Diff line change
Expand Up @@ -2438,10 +2438,6 @@ def reset_metadata(self) -> None:
self.kv_block_allocator.reset()
self.request_to_kv_block_ids.fill_(-1)

# Reset step counter and LRU clock
self.step_count = 0
self.prefix_cache_lru_clock = 0

# Reset chunked prefill state
self.chunked_prefill_request_id = -1
self.num_prefill_requests = 0
Expand All @@ -2466,6 +2462,11 @@ def reset(self) -> None:
self.reset_tensors()
self.reset_metadata()

# Reset lifetime counters (not reset in reset_metadata, which is also
# called during suspend/resume where these must persist).
self.step_count = 0
self.prefix_cache_lru_clock = 0

# Reset Mamba cache state
if self.mamba_slot_allocator is not None:
self.mamba_slot_allocator.reset()
Expand Down
97 changes: 62 additions & 35 deletions megatron/core/inference/engines/dynamic_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -241,8 +241,8 @@ def __init__(self, controller: TextGenerationController, context: DynamicInferen
if self.num_speculative_tokens > 0:
assert (
model_config.mtp_use_repeated_layer
or self.num_speculative_tokens <= self.controller.num_mtp_heads
), f"Number of speculative tokens {self.num_speculative_tokens} must be less than or equal to number of MTP heads {self.controller.num_mtp_heads}"
or self.num_speculative_tokens <= model_config.mtp_num_layers
), f"Number of speculative tokens {self.num_speculative_tokens} must be less than or equal to number of MTP layers {model_config.mtp_num_layers}"
self.track_paused_request_events = inference_config.track_paused_request_events
self.track_generated_token_events = inference_config.track_generated_token_events
self.enable_chunked_prefill = inference_config.enable_chunked_prefill
Expand Down Expand Up @@ -329,9 +329,15 @@ def reset(self) -> None:

self.resume_request_ids = None

# Speculative decoding acceptance tracking.
self._spec_tokens_proposed = 0
self._spec_tokens_accepted = 0
# Speculative decoding acceptance tracking (per-position).
# Each tensor has length num_speculative_tokens; index i tracks position i+1
# (i.e. the i-th draft token proposed by the MTP head).
self._spec_tokens_proposed_per_pos = torch.zeros(
self.num_speculative_tokens, dtype=torch.int64
)
self._spec_tokens_accepted_per_pos = torch.zeros(
self.num_speculative_tokens, dtype=torch.int64
)
self._spec_steps = 0

# Prefix caching tracking.
Expand Down Expand Up @@ -394,15 +400,15 @@ def create_cuda_graphs(self, reset_context: bool = True):
# 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
controller.num_mtp_depths > 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_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)
Expand Down Expand Up @@ -1200,6 +1206,7 @@ def post_process_requests(
tokens = accepted_tokens + tokens

num_stop_word_trim = 0
is_prefill = len(request.generated_tokens) == 0
if request_id != self.context.chunked_prefill_request_id:
# Skip appending token for requests being finished due to stop words
# (they already have their final token from the previous step)
Expand Down Expand Up @@ -1270,13 +1277,21 @@ def post_process_requests(
request
)

# Track acceptance statistics for logging.
if len(request.generated_tokens) > 0 and self.num_speculative_tokens > 0:
# Track per-position acceptance statistics for logging.
# Skip prefill requests: MTP heads only propose speculative tokens
# for decode requests, so counting prefill requests would inflate
# the denominator and artificially deflate the acceptance rate.
if (
not is_prefill
and len(request.generated_tokens) > 0
and self.num_speculative_tokens > 0
):
actual_proposed = max(0, self.num_speculative_tokens - num_stop_word_trim)
actual_accepted = max(0, len(accepted_tokens) - num_stop_word_trim)

self._spec_tokens_proposed += actual_proposed
self._spec_tokens_accepted += actual_accepted
self._spec_tokens_proposed_per_pos[:actual_proposed] += 1
accepted_t = torch.tensor(accepted_tokens_list[:actual_proposed])
self._spec_tokens_accepted_per_pos[:actual_proposed] += (
accepted_t != -1
).long()

if request_id in finished_request_ids:
# Reconstruct routing from per-block storage before popping.
Expand Down Expand Up @@ -1943,13 +1958,24 @@ async def async_bookkeep(
else:
metrics[f'inference/{key}'] = value

# Add speculative decoding acceptance metrics.
if self.num_speculative_tokens > 0 and self._spec_tokens_proposed > 0:
acceptance_rate = self._spec_tokens_accepted / self._spec_tokens_proposed
# Add speculative decoding acceptance metrics (aggregate + per-position).
total_proposed = sum(self._spec_tokens_proposed_per_pos)
total_accepted = sum(self._spec_tokens_accepted_per_pos)
if self.num_speculative_tokens > 0 and total_proposed > 0:
acceptance_rate = total_accepted / total_proposed
metrics['inference/spec_decode_acceptance_rate'] = float(acceptance_rate * 100.0)
metrics['inference/spec_decode_tokens_proposed'] = int(self._spec_tokens_proposed)
metrics['inference/spec_decode_tokens_accepted'] = int(self._spec_tokens_accepted)
metrics['inference/spec_decode_tokens_proposed'] = int(total_proposed)
metrics['inference/spec_decode_tokens_accepted'] = int(total_accepted)
metrics['inference/spec_decode_num_steps'] = int(self._spec_steps)
for pos in range(self.num_speculative_tokens):
if self._spec_tokens_proposed_per_pos[pos] > 0:
pos_rate = (
self._spec_tokens_accepted_per_pos[pos]
/ self._spec_tokens_proposed_per_pos[pos]
)
metrics[f'inference/spec_decode_acceptance_rate_pos{pos + 1}'] = float(
pos_rate * 100.0
)

# Add prefix caching metrics.
if self.context.enable_prefix_caching and self._prefix_cache_hits > 0:
Expand Down Expand Up @@ -2011,34 +2037,35 @@ async def async_bookkeep(
mem["reserved_bytes.all.current"] / (1024**3),
)
)
if self.num_speculative_tokens > 0 and self._spec_tokens_proposed > 0:
spec_rate = self._spec_tokens_accepted / self._spec_tokens_proposed * 100.0
output_str += " ... spec: accept %.1f%% (%d/%d in %d steps)" % (
total_proposed = sum(self._spec_tokens_proposed_per_pos)
total_accepted = sum(self._spec_tokens_accepted_per_pos)
if self.num_speculative_tokens > 0 and total_proposed > 0:
spec_rate = total_accepted / total_proposed * 100.0
per_pos_rates = []
for pos in range(self.num_speculative_tokens):
if self._spec_tokens_proposed_per_pos[pos] > 0:
pos_rate = (
self._spec_tokens_accepted_per_pos[pos]
/ self._spec_tokens_proposed_per_pos[pos]
* 100.0
)
per_pos_rates.append("t%d=%.1f%%" % (pos + 1, pos_rate))
output_str += " ... spec (cumul): accept %.1f%% (%d/%d in %d steps) [%s]" % (
spec_rate,
self._spec_tokens_accepted,
self._spec_tokens_proposed,
total_accepted,
total_proposed,
self._spec_steps,
", ".join(per_pos_rates),
)
if self.context.enable_prefix_caching and self._prefix_cache_hits > 0:
output_str += " ... prefix cache: %d hits, %d blocks matched" % (
output_str += " ... prefix cache (cumul): %d hits, %d blocks matched" % (
self._prefix_cache_hits,
self._prefix_cache_blocks_matched,
)
if context_state["is_decode_only"]:
output_str = f"\033[94m{output_str}\033[0m"
logging.info(output_str)

# Reset speculative decoding accumulators after both wandb and console logging.
if self.num_speculative_tokens > 0:
self._spec_tokens_proposed = 0
self._spec_tokens_accepted = 0
self._spec_steps = 0

# Reset prefix caching accumulators after both wandb and console logging.
if self.context.enable_prefix_caching:
self._prefix_cache_hits = 0
self._prefix_cache_blocks_matched = 0

nvtx_range_pop("console_logging")

return {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -104,9 +104,21 @@ def __init__(self, inference_wrapped_model: AbstractModelInferenceWrapper, token
self.vocab_size = unwrapped_model.vocab_size

self.sampling_rng = torch.Generator(device=torch.cuda.current_device())
self.num_mtp_heads = self._get_mtp_num_heads()
self.sampling_rng.manual_seed(self.model_config.inference_sampling_seed)

if not self.num_speculative_tokens:
self.num_mtp_depths = 0
else:
assert (
self.model_config.mtp_num_layers and self.model_config.mtp_num_layers >= 1
), "mtp_num_layers must be >= 1 when num_speculative_tokens > 0"
if self.model_config.mtp_use_repeated_layer:
self.num_mtp_depths = self.num_speculative_tokens
else:
self.num_mtp_depths = min(
self.num_speculative_tokens, self.model_config.mtp_num_layers
)

if (
self.model_config.cuda_graph_impl == "local"
and self.model_config.expert_model_parallel_size > 1
Expand All @@ -120,13 +132,6 @@ def __init__(self, inference_wrapped_model: AbstractModelInferenceWrapper, token
if self.inference_wrapped_model.inference_context.is_dynamic_batching():
self._init_dynamic_sampling_tensors()

def _get_mtp_num_heads(self) -> int:
"""Get the number of MTP layers from the model config."""
model = self.inference_wrapped_model.model
if hasattr(model, 'config') and hasattr(model.config, 'mtp_num_layers'):
return model.config.mtp_num_layers or 0
return 0

def set_stop_word_finished_ids_callback(self, callback):
"""Set a callback to get request IDs that should be marked as finished due to stop words.
Expand Down Expand Up @@ -222,7 +227,6 @@ def _init_mtp_sampling_tensors(self):
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
Expand Down Expand Up @@ -833,7 +837,7 @@ def _compute_serial_mtp_and_sample(self):
position_ids_buf[0, active_request_count:] = 0

nvtx_range_pop("mtp-spec-decoding/serial-mtp-init")
for depth in range(self._num_mtp_depths):
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
Expand Down Expand Up @@ -1508,7 +1512,7 @@ def _dummy_serial_mtp_forward(self):
- When PP > 1: participate in the ``broadcast_from_last_pipeline_stage``
that the real ranks also perform.
"""
if self.num_speculative_tokens == 0 or self.num_mtp_heads == 0:
if self.num_speculative_tokens == 0 or self.num_mtp_depths == 0:
return
if self.model_config.expert_model_parallel_size <= 1:
return
Expand Down Expand Up @@ -1549,7 +1553,7 @@ def _dummy_serial_mtp_forward(self):

context = self.inference_wrapped_model.inference_context

for depth in range(self._num_mtp_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:
Expand Down Expand Up @@ -1806,6 +1810,14 @@ async def async_generate_output_tokens_dynamic_batch(
)
range_pop()

# Capture before update_requests (called by _dynamic_step_context_bookkeeping)
# resets num_prefill_requests to 0, which would make num_decode_requests
# always equal to the full active count.
num_decode_requests = context.num_decode_requests
if self.num_speculative_tokens > 0:
# Prefill-only batches must not have any accepted speculative tokens.
assert num_decode_requests > 0 or (self._accepted_tokens_per_request == -1).all()

if skip_bookkeeping:
# _transfer_samples_to_cpu wasn't invoked on this path, so do
# a one-shot D2H here to keep "sample" as a CPU tensor for
Expand All @@ -1820,9 +1832,9 @@ async def async_generate_output_tokens_dynamic_batch(

ret = {
"accepted_tokens": (
# Clone needed: .fill_(-1) on line 1480 would corrupt the returned value.
# Clone needed: .fill_(-1) below would corrupt the returned value.
self._accepted_tokens_per_request.clone()
if self.num_speculative_tokens > 0
if self.num_speculative_tokens > 0 and num_decode_requests > 0
else None
),
"log_probs": log_probs,
Expand Down
Loading
Loading