Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
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
20 changes: 10 additions & 10 deletions megatron/core/models/gpt/gpt_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -665,18 +665,18 @@ def _postprocess(

if self.config.mtp_num_layers:
assert self.config.mtp_num_layers > 0
if in_inference_mode or is_spec_decode:
if is_spec_decode:
# Cache decoder hidden states for serial MTP computation
# after speculative token verification.
if inference_context is not None:
if self.config.inference_cuda_graph_scope == InferenceCudaGraphScope.block:
assert inference_context.mtp_decoder_hidden_states is not None
inference_context.mtp_decoder_hidden_states[: hidden_states.shape[0]].copy_(
hidden_states
)
else:
inference_context.mtp_decoder_hidden_states = hidden_states
else:
assert inference_context is not None
if self.config.inference_cuda_graph_scope == InferenceCudaGraphScope.block:
assert inference_context.mtp_decoder_hidden_states is not None
inference_context.mtp_decoder_hidden_states[: hidden_states.shape[0]].copy_(
hidden_states
)
else:
inference_context.mtp_decoder_hidden_states = hidden_states
elif not in_inference_mode:
# In training/eval, use the utility function for processing MTP loss/scaling.
hidden_states = process_mtp_loss(
hidden_states=hidden_states,
Expand Down
30 changes: 15 additions & 15 deletions megatron/core/models/hybrid/hybrid_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -543,21 +543,21 @@ def forward(

if self.config.mtp_num_layers is not None and self.mtp_process:
assert self.config.mtp_num_layers > 0
if in_inference_mode or is_spec_decode:
if inference_context is not None:
if self.config.inference_cuda_graph_scope == InferenceCudaGraphScope.block:
# Block-scope CUDA graph mode: copy_() into the
# pre-allocated buffer so every graph replay writes to
# the same fixed GPU address regardless of batch size.
assert inference_context.mtp_decoder_hidden_states is not None
inference_context.mtp_decoder_hidden_states[: hidden_states.shape[0]].copy_(
hidden_states
)
else:
# Non-block scope: direct assignment; the controller will set
# this back to None after reading to allow GC.
inference_context.mtp_decoder_hidden_states = hidden_states
else:
if is_spec_decode:
assert inference_context is not None
if self.config.inference_cuda_graph_scope == InferenceCudaGraphScope.block:
# Block-scope CUDA graph mode: copy_() into the
# pre-allocated buffer so every graph replay writes to
# the same fixed GPU address regardless of batch size.
assert inference_context.mtp_decoder_hidden_states is not None
inference_context.mtp_decoder_hidden_states[: hidden_states.shape[0]].copy_(
hidden_states
)
else:
# Non-block scope: direct assignment; the controller will set
# this back to None after reading to allow GC.
inference_context.mtp_decoder_hidden_states = hidden_states
elif not in_inference_mode:
# For RL (labels is None), process_mtp_loss derives labels from
# input_ids to match the SFT label format.
hidden_states = process_mtp_loss(
Expand Down
128 changes: 108 additions & 20 deletions tests/unit_tests/inference/test_mtp_cuda_graph_inference.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,6 +33,7 @@
from megatron.core.inference.text_generation_controllers.text_generation_controller import (
TextGenerationController,
)
from megatron.core.inference.utils import InferenceMode
from megatron.core.models.backends import LocalSpecProvider
from megatron.core.models.gpt.gpt_layer_specs import (
get_gpt_layer_local_spec,
Expand Down Expand Up @@ -1174,8 +1175,8 @@ def teardown_class(cls):
def teardown_method(self):
delete_cuda_graphs()

def _build_model(self, *, inference_cuda_graph_scope='block'):
"""Build a HybridModel with MTP and block-scope CUDA graph support."""
def _build_model(self, *, inference_cuda_graph_scope='block', model_type='hybrid'):
"""Build a GPT or Hybrid model with MTP 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,
Expand All @@ -1191,35 +1192,58 @@ def _build_model(self, *, inference_cuda_graph_scope='block'):
cuda_graph_impl="local",
inference_cuda_graph_scope=inference_cuda_graph_scope,
)
hybrid_stack_spec = _build_hybrid_stack_spec()
model = HybridModel(
config=config,
hybrid_stack_spec=hybrid_stack_spec,
vocab_size=self.VOCAB_SIZE,
max_sequence_length=self.MAX_SEQ_LEN,
parallel_output=True,
pre_process=True,
post_process=True,
hybrid_layer_pattern="****/*",
position_embedding_type='rope',
).cuda()
if model_type == 'gpt':
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,
mtp_block_spec=mtp_block_spec,
vocab_size=self.VOCAB_SIZE,
max_sequence_length=self.MAX_SEQ_LEN,
parallel_output=True,
pre_process=True,
post_process=True,
position_embedding_type='rope',
).cuda()
elif model_type == 'hybrid':
hybrid_stack_spec = _build_hybrid_stack_spec()
model = HybridModel(
config=config,
hybrid_stack_spec=hybrid_stack_spec,
vocab_size=self.VOCAB_SIZE,
max_sequence_length=self.MAX_SEQ_LEN,
parallel_output=True,
pre_process=True,
post_process=True,
hybrid_layer_pattern="****/*",
position_embedding_type='rope',
).cuda()
else:
raise ValueError(f"Unknown model_type: {model_type!r}")
for param in model.parameters():
param.data = param.data.to(config.params_dtype)
model.eval()
return model

def _build_engine(self, *, inference_cuda_graph_scope='block'):
def _build_engine(
self, *, inference_cuda_graph_scope='block', num_speculative_tokens=1, model_type='hybrid'
):
"""Build a DynamicInferenceEngine with block-scope CUDA graphs."""
delete_cuda_graphs()
model = self._build_model(inference_cuda_graph_scope=inference_cuda_graph_scope)
model = self._build_model(
inference_cuda_graph_scope=inference_cuda_graph_scope, model_type=model_type
)
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=1,
num_speculative_tokens=num_speculative_tokens,
block_size_tokens=256,
max_requests=16,
num_cuda_graphs=-1,
Expand All @@ -1233,18 +1257,21 @@ def _build_engine(self, *, inference_cuda_graph_scope='block'):
engine = DynamicInferenceEngine(ctrl, context)
return engine

@pytest.mark.parametrize("model_type", ['gpt', 'hybrid'])
@pytest.mark.parametrize("inference_cuda_graph_scope", ['block', 'layer'])
@torch.inference_mode()
def test_decoder_hidden_states_set_after_forward(self, inference_cuda_graph_scope):
def test_decoder_hidden_states_set_after_forward(self, inference_cuda_graph_scope, model_type):
"""Decoder hidden states are accessible via the context after each forward pass.

Block-scope CUDA graphs: forward() writes via copy_() into the pre-allocated
context buffer, captured once and replayed to the same GPU address each step.
Layer-scope (non-block) CUDA graphs: forward() assigns the tensor directly to
the context attribute; the controller sets it back to None after reading to allow GC.
Both scopes are valid with cuda_graph_impl='local'.
Both scopes are valid with cuda_graph_impl='local'. Covers GPTModel and HybridModel.
"""
engine = self._build_engine(inference_cuda_graph_scope=inference_cuda_graph_scope)
engine = self._build_engine(
inference_cuda_graph_scope=inference_cuda_graph_scope, model_type=model_type
)
ctrl = engine.controller
context = engine.context

Expand Down Expand Up @@ -1382,3 +1409,64 @@ def _run_eager_mtp(decoder_hidden_states):
f"{sampled.tolist()} != reference {reference_tokens[depth].tolist()}; "
"the unused buffer tail leaked into the MTP forward"
)

@pytest.mark.parametrize("model_type", ['gpt', 'hybrid'])
@pytest.mark.parametrize("inference_cuda_graph_scope", ['block', 'layer'])
@torch.inference_mode()
def test_no_spec_decode_leaves_decoder_hidden_states_unset(
self, inference_cuda_graph_scope, model_type
):
"""Regression: a model with an MTP head but ``num_speculative_tokens == 0``.

When the model has MTP layers (``mtp_num_layers >= 1``) but speculative
decoding is disabled, plain inference must NOT touch
``context.mtp_decoder_hidden_states`` — there is no serial post-verification
MTP step to consume it, and for block-scope CUDA graphs the buffer is never
even allocated (it is allocated only when ``num_speculative_tokens > 0``).

Covers both GPTModel and HybridModel since each carries the same MTP
post-process branch (``gpt_model.py`` / ``hybrid_model.py``).
"""
engine = self._build_engine(
inference_cuda_graph_scope=inference_cuda_graph_scope,
num_speculative_tokens=0,
model_type=model_type,
)
ctrl = engine.controller
context = engine.context

# No speculative decoding -> no MTP depths and no pre-allocated buffer.
assert ctrl.num_speculative_tokens == 0
assert ctrl.num_mtp_depths == 0
assert context.mtp_decoder_hidden_states is None

prompt_length = 10
req = DynamicInferenceRequest(
request_id=0,
prompt_tokens=torch.arange(prompt_length, device='cuda'),
sampling_params=SamplingParams(num_tokens_to_generate=20),
)
context.add_request(req)
context.initialize_attention_state()

active_mask = torch.ones(1, device='cuda', dtype=torch.int32)
new_tokens = torch.zeros(1, device='cuda', dtype=torch.int64)
context.update_requests(
active_requests_mask=active_mask, new_tokens=new_tokens, new_speculative_tokens=None
)
context.initialize_attention_state()

# Force the inference flag on so the forward takes the in_inference_mode
# branch even though we drive the step directly rather than via the engine
# run loop.
with InferenceMode.active():
for step in range(3):
input_ids, position_ids = ctrl._dynamic_step_context_init()
ctrl._dynamic_step_forward_logits(input_ids, position_ids)

assert context.mtp_decoder_hidden_states is None, (
f"Step {step}: mtp_decoder_hidden_states should stay None when "
f"num_speculative_tokens == 0 (model={model_type}, "
f"scope={inference_cuda_graph_scope}), got a tensor of shape "
f"{tuple(context.mtp_decoder_hidden_states.shape)}"
)
Loading