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
2 changes: 1 addition & 1 deletion 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
5 changes: 4 additions & 1 deletion megatron/training/training.py
Original file line number Diff line number Diff line change
Expand Up @@ -588,6 +588,9 @@ def transformer_flops():
# Calculate the number of each type of layer.
num_attn_layers, num_mamba_layers, num_mlp_layers, num_moe_layers = calculate_layer_counts()

mtp_num_layers = args.mtp_num_layers
if mtp_num_layers is None:
mtp_num_layers = 0
# Compute hybrid model FLOPs.
return hybrid_flops(
batch_size=batch_size,
Expand All @@ -614,7 +617,7 @@ def transformer_flops():
else args.moe_shared_expert_intermediate_size),
num_experts_routed_to=args.moe_router_topk,
vocab_size=args.padded_vocab_size,
mtp_num_layers=args.mtp_num_layers,
mtp_num_layers=mtp_num_layers,
)
else:
# Compute standard Transformer model FLOPs.
Expand Down
Loading