Skip to content
Open
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
2 changes: 1 addition & 1 deletion megatron/core/models/common/model_chunk_schedule_plan.py
Original file line number Diff line number Diff line change
Expand Up @@ -123,7 +123,7 @@ def _build_callable_nodes(self, event, comp_stream, comm_stream, extra_args):

# get flags for latter use
is_mtp = isinstance(self.layer, MultiTokenPredictionLayer)
transformer_layer = self.layer.transformer_layer if is_mtp else self.layer
transformer_layer = self.layer.mtp_model_layer if is_mtp else self.layer
is_moe = isinstance(transformer_layer.mlp, MoELayer)
num_local_experts = transformer_layer.mlp.num_local_experts if is_moe else None

Expand Down
4 changes: 2 additions & 2 deletions megatron/core/models/gpt/fine_grained_callables.py
Original file line number Diff line number Diff line change
Expand Up @@ -638,9 +638,9 @@ def build_mtp_layer_callables(layer):
multi-token prediction layer nodes (attention, MLP, etc.)
"""

forward_funcs, backward_dw = build_transformer_layer_callables(layer.transformer_layer)
forward_funcs, backward_dw = build_transformer_layer_callables(layer.mtp_model_layer)
attn_forward, dispatch_forward, mlp_forward, combine_forward, _ = forward_funcs
is_moe = isinstance(layer.transformer_layer.mlp, MoELayer)
is_moe = isinstance(layer.mtp_model_layer.mlp, MoELayer)
assert is_moe, "MTP layer in a2a overlap only supports MoE layer for now."

def submodule_mtp_attn_forward(node, hidden_states):
Expand Down
2 changes: 1 addition & 1 deletion megatron/core/models/gpt/gpt_layer_specs.py
Original file line number Diff line number Diff line change
Expand Up @@ -753,7 +753,7 @@ def get_gpt_mtp_block_spec_for_backend(
raise ValueError(f"Invalid spec: {spec}")

mtp_layer_spec = get_mtp_layer_spec_for_backend(
transformer_layer_spec=transformer_layer_spec, backend=backend
mtp_model_layer_spec=transformer_layer_spec, backend=backend
)
mtp_num_layers = config.mtp_num_layers if config.mtp_num_layers else 0
mtp_layer_specs = [mtp_layer_spec] * mtp_num_layers
Expand Down
2 changes: 1 addition & 1 deletion megatron/core/pipeline_parallel/schedules.py
Original file line number Diff line number Diff line change
Expand Up @@ -213,7 +213,7 @@ def set_current_microbatch(model, microbatch_id):
layer.current_microbatch = microbatch_id
if hasattr(model_with_decoder, 'mtp'):
for layer in model_with_decoder.mtp.layers:
layer.transformer_layer.current_microbatch = microbatch_id
layer.mtp_model_layer.current_microbatch = microbatch_id


def forward_step_calc_loss(
Expand Down
4 changes: 2 additions & 2 deletions megatron/core/transformer/cuda_graphs.py
Original file line number Diff line number Diff line change
Expand Up @@ -1738,7 +1738,7 @@ def __init__(self, model, config, seq_length, micro_batch_size, optimizers=[]):
callables.append(layer)
callables_is_mtp.append(False)
for layer_number in range(num_mtp_layers):
layer = chunk_with_decoder.mtp.layers[layer_number].transformer_layer
layer = chunk_with_decoder.mtp.layers[layer_number].mtp_model_layer
if _layer_is_graphable(layer, config):
num_graphable_layers += 1
callables.append(layer)
Expand Down Expand Up @@ -1855,7 +1855,7 @@ def _get_layer_static_inputs(layer, chunk_of_the_layer):
Get the static inputs for a layer.
"""
assert layer in chunk_of_the_layer.decoder.layers or any(
layer is mtp_layer.transformer_layer for mtp_layer in chunk_of_the_layer.mtp.layers
layer is mtp_layer.mtp_model_layer for mtp_layer in chunk_of_the_layer.mtp.layers
), "Layer is not in the chunk"

def get_rotary_pos_emb(transformer_module, transformer_input):
Expand Down
24 changes: 12 additions & 12 deletions megatron/core/transformer/multi_token_prediction.py
Original file line number Diff line number Diff line change
Expand Up @@ -397,33 +397,33 @@ class MultiTokenPredictionLayerSubmodules:
embedding normalization to be applied.
eh_proj (Union[ModuleSpec, type]): Specification or instance of the
linear projection to be applied.
transformer_layer (Union[ModuleSpec, type]): Specification
mtp_model_layer (Union[ModuleSpec, type]): Specification
or instance of the transformer block to be applied.
"""

enorm: Union[ModuleSpec, type] = None
hnorm: Union[ModuleSpec, type] = None
eh_proj: Union[ModuleSpec, type] = None
transformer_layer: Union[ModuleSpec, type] = None
mtp_model_layer: Union[ModuleSpec, type] = None
layer_norm: Union[ModuleSpec, type] = None


def get_mtp_layer_spec(
transformer_layer_spec: ModuleSpec, use_transformer_engine: bool
mtp_model_layer_spec: ModuleSpec, use_transformer_engine: bool
) -> ModuleSpec:
"""Get the MTP layer spec.

Returns:
ModuleSpec: Module specification with TE modules
"""
return get_mtp_layer_spec_for_backend(
transformer_layer_spec,
mtp_model_layer_spec,
backend=TESpecProvider() if use_transformer_engine else LocalSpecProvider(),
)


def get_mtp_layer_spec_for_backend(
transformer_layer_spec: ModuleSpec, backend: BackendSpecProvider
mtp_model_layer_spec: ModuleSpec, backend: BackendSpecProvider
) -> ModuleSpec:
"""Get the MTP layer spec.

Expand All @@ -438,7 +438,7 @@ def get_mtp_layer_spec_for_backend(
enorm=layer_norm_impl,
hnorm=layer_norm_impl,
eh_proj=column_parallel_linear_impl,
transformer_layer=transformer_layer_spec,
mtp_model_layer=mtp_model_layer_spec,
layer_norm=layer_norm_impl,
),
)
Expand Down Expand Up @@ -622,7 +622,7 @@ def __init__(
self.vp_stage = vp_stage
self.cp_group = pg_collection.cp

self_attention_spec = self.submodules.transformer_layer.submodules.self_attention
self_attention_spec = self.submodules.mtp_model_layer.submodules.self_attention
attn_mask_type = self_attention_spec.params.get('attn_mask_type', '')
assert attn_mask_type in SUPPORTED_ATTN_MASK, (
f"Multi-Token Prediction (MTP) is not jet supported with "
Expand Down Expand Up @@ -664,14 +664,14 @@ def __init__(
diff_transformer_layer_offset = self.config.num_layers - get_transformer_layer_offset(
self.config, vp_stage
)
self.transformer_layer = build_module(
self.submodules.transformer_layer,
self.mtp_model_layer = build_module(
self.submodules.mtp_model_layer,
config=self.config,
vp_stage=vp_stage,
layer_number=self.layer_number + diff_transformer_layer_offset,
)
if hasattr(self.transformer_layer.mlp, 'set_is_mtp'):
self.transformer_layer.mlp.set_is_mtp()
if hasattr(self.mtp_model_layer.mlp, 'set_is_mtp'):
self.mtp_model_layer.mlp.set_is_mtp()

self.final_layernorm = build_module(
self.submodules.layer_norm,
Expand Down Expand Up @@ -793,7 +793,7 @@ def _proj_and_transformer_layer(
# transformer layer is cudagraphed, the FP8GlobalStateManager.is_first_fp8_module() is
# True so that the fp8 weight caching can be triggered correctly.
with transformer_layer_fp8_context:
hidden_states, _ = self.transformer_layer(
hidden_states, _ = self.mtp_model_layer(
hidden_states=hidden_states,
attention_mask=attention_mask,
context=context,
Expand Down